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
6use memra_gguf::config::AttentionGateKind;
7use memra_gguf::model_plan::{
8    ActivationPlan, AttentionPlan, AttentionScale, GemmaLayerScale, LogitsTransform, MlpPlan,
9    ModelPlan, ResidualTopology, ValueNorm, ValueProjection,
10};
11use memra_gguf::tensor_contract::{DsparkTensor, LayerTensor, MtpTensor, TensorId, VisionTensor};
12use std::collections::BTreeMap;
13
14#[derive(Debug, Clone, PartialEq)]
15pub struct ReferenceTensor {
16    /// Logical row-major shape, outermost dimension first.
17    pub shape: Vec<usize>,
18    pub data: Vec<f32>,
19}
20
21impl ReferenceTensor {
22    pub fn new(shape: Vec<usize>, data: Vec<f32>) -> Result<Self, ReferenceError> {
23        let expected = shape.iter().product();
24        if data.len() != expected {
25            return Err(ReferenceError::TensorShape {
26                id: None,
27                expected: shape,
28                actual_elements: data.len(),
29            });
30        }
31        Ok(Self { shape, data })
32    }
33}
34
35pub type ReferenceWeights = BTreeMap<TensorId, ReferenceTensor>;
36
37#[derive(Debug, Clone, PartialEq)]
38pub struct ReferenceFixture {
39    pub token_ids: Vec<u32>,
40    pub weights: ReferenceWeights,
41    pub vision: Option<ReferenceVisionInput>,
42    pub multimodal_token_ids: Option<Vec<u32>>,
43}
44
45#[derive(Debug, Clone, PartialEq)]
46pub struct ReferenceVisionInput {
47    /// Raw patch pixels in `[0, 1]`, row-major `[patches, 3 * patch_size^2]`.
48    pub patches: ReferenceTensor,
49    pub positions: Vec<[u32; 2]>,
50    pub output_tokens: usize,
51}
52
53#[derive(Debug, Clone, PartialEq)]
54pub struct ReferenceVisionOutput {
55    pub encoder_hidden: Vec<f32>,
56    pub pooled_hidden: Vec<f32>,
57    pub projected_hidden: Vec<f32>,
58    pub patch_count: usize,
59    pub output_tokens: usize,
60    pub hidden_size: usize,
61    pub projection_size: usize,
62}
63
64#[derive(Debug, Clone, PartialEq)]
65pub struct ReferenceMultimodalOutput {
66    pub language: ReferenceOutput,
67    pub vision: ReferenceVisionOutput,
68}
69
70#[derive(Debug, Clone, PartialEq)]
71pub struct ReferenceState {
72    pub layers: Vec<ReferenceLayerState>,
73}
74
75#[derive(Debug, Clone, PartialEq)]
76pub enum ReferenceLayerState {
77    Kv {
78        key: Vec<f32>,
79        value: Vec<f32>,
80        tokens: usize,
81        kv_heads: usize,
82        key_head_dim: usize,
83        value_head_dim: usize,
84        window: Option<usize>,
85    },
86    Recurrent {
87        conv: Vec<f32>,
88        matrix: Vec<f32>,
89        value_heads: usize,
90        key_head_dim: usize,
91        value_head_dim: usize,
92        conv_width: usize,
93    },
94    LatentKv {
95        rows: Vec<f32>,
96        tokens: usize,
97        width: usize,
98    },
99    CompressedAttention {
100        rows: Vec<f32>,
101        tokens: usize,
102        width: usize,
103        window: usize,
104        compressed_tokens: usize,
105    },
106}
107
108#[derive(Debug, Clone, PartialEq)]
109pub struct ReferenceOutput {
110    /// `[tokens, vocab]`, row-major.
111    pub logits: Vec<f32>,
112    pub tokens: usize,
113    pub vocab: usize,
114    pub state: ReferenceState,
115    pub mtp: Vec<ReferenceMtpOutput>,
116    pub draft: Option<ReferenceDraftOutput>,
117}
118
119#[derive(Debug, Clone, PartialEq)]
120pub struct ReferenceMtpOutput {
121    pub depth: u32,
122    pub logits: Vec<f32>,
123    pub hidden: Vec<f32>,
124    pub state: ReferenceLayerState,
125}
126
127#[derive(Debug, Clone, PartialEq)]
128pub struct ReferenceDraftOutput {
129    pub input_token: u32,
130    pub output_ids: Vec<u32>,
131    pub confidence: Vec<f32>,
132    pub logits: Vec<f32>,
133    pub hidden: Vec<f32>,
134    pub block_size: usize,
135}
136
137#[derive(Debug, Clone, PartialEq)]
138pub enum ReferenceError {
139    EmptyInput,
140    TokenOutOfRange {
141        token: u32,
142        vocab: usize,
143    },
144    MissingTensor(TensorId),
145    TensorShape {
146        id: Option<TensorId>,
147        expected: Vec<usize>,
148        actual_elements: usize,
149    },
150    UnsupportedOperation {
151        layer: Option<u32>,
152        operation: &'static str,
153    },
154    InvalidPlan {
155        layer: Option<u32>,
156        reason: &'static str,
157    },
158}
159
160impl std::fmt::Display for ReferenceError {
161    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
162        match self {
163            Self::EmptyInput => write!(f, "reference executor requires at least one token"),
164            Self::TokenOutOfRange { token, vocab } => {
165                write!(f, "token {token} is outside vocabulary size {vocab}")
166            }
167            Self::MissingTensor(id) => write!(f, "missing reference tensor {id:?}"),
168            Self::TensorShape {
169                id,
170                expected,
171                actual_elements,
172            } => write!(
173                f,
174                "reference tensor {id:?} expected shape {expected:?}, got {actual_elements} elements"
175            ),
176            Self::UnsupportedOperation { layer, operation } => {
177                write!(
178                    f,
179                    "unsupported reference operation {operation} at layer {layer:?}"
180                )
181            }
182            Self::InvalidPlan { layer, reason } => {
183                write!(f, "invalid model plan at layer {layer:?}: {reason}")
184            }
185        }
186    }
187}
188
189impl std::error::Error for ReferenceError {}
190
191pub fn deterministic_fixture(plan: &ModelPlan) -> Result<ReferenceFixture, ReferenceError> {
192    let hidden = plan.hidden_size as usize;
193    let vocab = plan.vocab_size as usize;
194    if hidden == 0 || vocab < 2 || hidden > 256 || vocab > 262_144 {
195        return Err(ReferenceError::InvalidPlan {
196            layer: None,
197            reason: "reference fixture requires hidden<=256 and 2<=vocab<=262144",
198        });
199    }
200    let mut executable_layers: Vec<_> = plan
201        .layers
202        .iter()
203        .chain(plan.mtp_blocks.iter().map(|block| &block.layer))
204        .collect();
205    if let Some(memra_gguf::model_plan::DrafterPlan::Dspark(dspark)) = plan.drafter.as_ref() {
206        executable_layers.extend(dspark.blocks.iter());
207    }
208    let mut weights = ReferenceWeights::new();
209    weights.insert(
210        TensorId::TokenEmbedding,
211        generated_tensor(&[vocab, hidden], 1, 0.2)?,
212    );
213    let vision = if let Some(vision) = plan.vision.as_ref() {
214        Some(add_vision_fixture(
215            &mut weights,
216            vision,
217            plan.multimodal.map(|injection| injection.tokens_per_image),
218        )?)
219    } else {
220        None
221    };
222    weights.insert(
223        TensorId::OutputNorm,
224        ReferenceTensor::new(vec![hidden], vec![1.0; hidden])?,
225    );
226    let checkpoint_factor_width = executable_layers
227        .iter()
228        .copied()
229        .filter_map(|layer| match &layer.attention {
230            AttentionPlan::Full(attention) | AttentionPlan::SlidingWindow { attention, .. } => {
231                matches!(
232                    attention.rope.factors,
233                    memra_gguf::model_plan::RopeFactors::Checkpoint
234                )
235                .then_some(attention.rope.dimensions as usize / 2)
236            }
237            _ => None,
238        })
239        .max();
240    if let Some(width) = checkpoint_factor_width {
241        weights.insert(
242            TensorId::RopeFactors,
243            ReferenceTensor::new(vec![width], vec![1.0; width])?,
244        );
245    }
246    if let Some((streams, epsilon, sinkhorn_iterations)) = hyper_topology(plan)? {
247        add_hyper_head_fixture(&mut weights, streams, hidden)?;
248        if epsilon <= 0.0 || sinkhorn_iterations == 0 {
249            return Err(ReferenceError::InvalidPlan {
250                layer: None,
251                reason: "HyperConnections require positive epsilon and Sinkhorn iterations",
252            });
253        }
254    }
255    for layer in executable_layers {
256        match layer.residual {
257            ResidualTopology::Serial => {}
258            ResidualTopology::Gemma { parallel_moe, .. } => {
259                for tensor in [LayerTensor::PostAttentionNorm, LayerTensor::PostMlpNorm] {
260                    weights.insert(
261                        layer_id(layer.index, tensor),
262                        ReferenceTensor::new(vec![hidden], vec![1.0; hidden])?,
263                    );
264                }
265                weights.insert(
266                    layer_id(layer.index, LayerTensor::LayerScale),
267                    ReferenceTensor::new(vec![1], vec![0.9])?,
268                );
269                if parallel_moe.is_some() {
270                    for tensor in [
271                        LayerTensor::PostSharedMlpNorm,
272                        LayerTensor::PreRoutedMlpNorm,
273                        LayerTensor::PostRoutedMlpNorm,
274                    ] {
275                        weights.insert(
276                            layer_id(layer.index, tensor),
277                            ReferenceTensor::new(vec![hidden], vec![1.0; hidden])?,
278                        );
279                    }
280                }
281            }
282            ResidualTopology::HyperConnections { streams, .. } => {
283                add_hyper_fixture(&mut weights, layer.index, streams as usize, hidden)?;
284            }
285        }
286        for tensor in [LayerTensor::PreAttentionNorm, LayerTensor::PreMlpNorm] {
287            weights.insert(
288                layer_id(layer.index, tensor),
289                ReferenceTensor::new(vec![hidden], vec![1.0; hidden])?,
290            );
291        }
292        match &layer.attention {
293            AttentionPlan::Full(attention) | AttentionPlan::SlidingWindow { attention, .. } => {
294                add_full_attention_fixture(&mut weights, layer.index, attention, hidden)?;
295            }
296            AttentionPlan::GatedDeltaNet(gdn) => {
297                add_gdn_fixture(&mut weights, layer.index, gdn, hidden)?;
298            }
299            AttentionPlan::Mla(mla) => {
300                add_mla_fixture(&mut weights, layer.index, mla, hidden)?;
301            }
302        }
303        match &layer.mlp {
304            MlpPlan::Dense(mlp) => {
305                add_dense_mlp_fixture(&mut weights, layer.index, mlp, hidden)?;
306            }
307            MlpPlan::Moe(moe) => {
308                add_moe_fixture(&mut weights, layer.index, moe, hidden, vocab)?;
309                if matches!(
310                    layer.residual,
311                    ResidualTopology::Gemma {
312                        parallel_moe: Some(_),
313                        ..
314                    }
315                ) {
316                    add_gemma_parallel_moe_fixture(&mut weights, layer.index, moe, hidden)?;
317                }
318            }
319        }
320    }
321    for block in &plan.mtp_blocks {
322        for tensor in [MtpTensor::EmbeddingNorm, MtpTensor::HiddenNorm] {
323            weights.insert(
324                TensorId::Mtp {
325                    depth: block.depth,
326                    tensor,
327                },
328                ReferenceTensor::new(vec![hidden], vec![1.0; hidden])?,
329            );
330        }
331        weights.insert(
332            TensorId::Mtp {
333                depth: block.depth,
334                tensor: MtpTensor::FusionProjection,
335            },
336            generated_tensor(
337                &[hidden, 2 * hidden],
338                100 + block.depth as u64,
339                1.0 / ((2 * hidden) as f32).sqrt(),
340            )?,
341        );
342    }
343    if let Some(memra_gguf::model_plan::DrafterPlan::Dspark(dspark)) = plan.drafter.as_ref() {
344        add_dspark_fixture(&mut weights, dspark, hidden, vocab)?;
345    }
346    let token_ids = (1..=3.min(vocab - 1)).map(|token| token as u32).collect();
347    let multimodal_token_ids = plan.multimodal.map(|injection| {
348        let mut tokens = Vec::with_capacity(injection.tokens_per_image as usize + 2);
349        tokens.push(1);
350        tokens.extend(std::iter::repeat_n(
351            injection.placeholder_token_id,
352            injection.tokens_per_image as usize,
353        ));
354        tokens.push(if injection.placeholder_token_id == 2 {
355            3
356        } else {
357            2
358        });
359        tokens
360    });
361    Ok(ReferenceFixture {
362        token_ids,
363        weights,
364        vision,
365        multimodal_token_ids,
366    })
367}
368
369fn add_dspark_fixture(
370    weights: &mut ReferenceWeights,
371    plan: &memra_gguf::model_plan::DsparkPlan,
372    hidden: usize,
373    vocab: usize,
374) -> Result<(), ReferenceError> {
375    if plan.blocks.is_empty()
376        || plan.block_size == 0
377        || plan.markov_rank == 0
378        || plan.target_layer_ids.is_empty()
379        || plan.noise_token_id as usize >= vocab
380    {
381        return Err(ReferenceError::InvalidPlan {
382            layer: None,
383            reason: "DSpark fixture requires blocks, targets, rank, block size, and valid noise token",
384        });
385    }
386    let streams = match plan.blocks[0].residual {
387        ResidualTopology::HyperConnections { streams, .. } if streams > 0 => streams as usize,
388        _ => {
389            return Err(ReferenceError::InvalidPlan {
390                layer: Some(plan.blocks[0].index),
391                reason: "DSpark blocks require HyperConnections",
392            });
393        }
394    };
395    weights.insert(
396        TensorId::Dspark(DsparkTensor::MainProjection),
397        generated_tensor(
398            &[hidden, plan.target_layer_ids.len() * hidden],
399            140,
400            1.0 / ((plan.target_layer_ids.len() * hidden) as f32).sqrt(),
401        )?,
402    );
403    weights.insert(
404        TensorId::Dspark(DsparkTensor::MainNorm),
405        ReferenceTensor::new(vec![hidden], vec![1.0; hidden])?,
406    );
407    weights.insert(
408        TensorId::Dspark(DsparkTensor::OutputNorm),
409        ReferenceTensor::new(vec![hidden], vec![1.0; hidden])?,
410    );
411    let rank = plan.markov_rank as usize;
412    weights.insert(
413        TensorId::Dspark(DsparkTensor::MarkovEmbedding),
414        generated_tensor(&[vocab, rank], 141, 0.1)?,
415    );
416    weights.insert(
417        TensorId::Dspark(DsparkTensor::MarkovOutput),
418        generated_tensor(&[vocab, rank], 142, 0.1)?,
419    );
420    weights.insert(
421        TensorId::Dspark(DsparkTensor::ConfidenceProjection),
422        generated_tensor(&[1, hidden + rank], 143, 0.1)?,
423    );
424    weights.insert(
425        TensorId::Dspark(DsparkTensor::HeadHyperFunction),
426        generated_tensor(&[streams, streams * hidden], 144, 0.1)?,
427    );
428    weights.insert(
429        TensorId::Dspark(DsparkTensor::HeadHyperBase),
430        generated_tensor(&[streams], 145, 0.05)?,
431    );
432    weights.insert(
433        TensorId::Dspark(DsparkTensor::HeadHyperScale),
434        ReferenceTensor::new(vec![1], vec![0.2])?,
435    );
436    Ok(())
437}
438
439fn add_vision_fixture(
440    weights: &mut ReferenceWeights,
441    plan: &memra_gguf::model_plan::VisionEncoderPlan,
442    output_tokens: Option<u32>,
443) -> Result<ReferenceVisionInput, ReferenceError> {
444    let hidden = plan.hidden_size as usize;
445    let patch_width =
446        (plan.patch.channels * plan.patch.patch_size * plan.patch.patch_size) as usize;
447    let axes = plan.patch.position_axes as usize;
448    let positions = plan.patch.position_embedding_size as usize;
449    weights.insert(
450        TensorId::Vision {
451            layer: None,
452            tensor: VisionTensor::PatchProjection,
453        },
454        generated_tensor(
455            &[hidden, patch_width],
456            150,
457            1.0 / (patch_width as f32).sqrt(),
458        )?,
459    );
460    weights.insert(
461        TensorId::Vision {
462            layer: None,
463            tensor: VisionTensor::PositionEmbedding,
464        },
465        generated_tensor(&[axes, positions, hidden], 151, 0.05)?,
466    );
467    if plan.standardize {
468        weights.insert(
469            TensorId::Vision {
470                layer: None,
471                tensor: VisionTensor::StandardizeBias,
472            },
473            generated_tensor(&[hidden], 152, 0.05)?,
474        );
475        weights.insert(
476            TensorId::Vision {
477                layer: None,
478                tensor: VisionTensor::StandardizeScale,
479            },
480            ReferenceTensor::new(vec![hidden], vec![0.5; hidden])?,
481        );
482    }
483    weights.insert(
484        TensorId::Vision {
485            layer: None,
486            tensor: VisionTensor::OutputProjection,
487        },
488        generated_tensor(
489            &[plan.projection_output_size as usize, hidden],
490            153,
491            1.0 / (hidden as f32).sqrt(),
492        )?,
493    );
494    for layer in &plan.layers {
495        let layer_id = Some(layer.index);
496        for tensor in [
497            VisionTensor::InputNorm,
498            VisionTensor::PostAttentionNorm,
499            VisionTensor::PreMlpNorm,
500            VisionTensor::PostMlpNorm,
501        ] {
502            weights.insert(
503                TensorId::Vision {
504                    layer: layer_id,
505                    tensor,
506                },
507                ReferenceTensor::new(vec![hidden], vec![1.0; hidden])?,
508            );
509        }
510        let query_width = (layer.attention.query_heads * layer.attention.head_dim) as usize;
511        let kv_width = (layer.attention.kv_heads * layer.attention.head_dim) as usize;
512        for (tensor, shape, input, salt) in [
513            (VisionTensor::Query, vec![query_width, hidden], hidden, 160),
514            (VisionTensor::Key, vec![kv_width, hidden], hidden, 161),
515            (VisionTensor::Value, vec![kv_width, hidden], hidden, 162),
516            (
517                VisionTensor::AttentionOutput,
518                vec![hidden, query_width],
519                query_width,
520                163,
521            ),
522            (
523                VisionTensor::MlpGate,
524                vec![layer.mlp.intermediate_size as usize, hidden],
525                hidden,
526                164,
527            ),
528            (
529                VisionTensor::MlpUp,
530                vec![layer.mlp.intermediate_size as usize, hidden],
531                hidden,
532                165,
533            ),
534            (
535                VisionTensor::MlpDown,
536                vec![hidden, layer.mlp.intermediate_size as usize],
537                layer.mlp.intermediate_size as usize,
538                166,
539            ),
540        ] {
541            weights.insert(
542                TensorId::Vision {
543                    layer: layer_id,
544                    tensor,
545                },
546                generated_tensor(
547                    &shape,
548                    salt + layer.index as u64 * 17,
549                    1.0 / (input as f32).sqrt(),
550                )?,
551            );
552        }
553        for tensor in [VisionTensor::QueryNorm, VisionTensor::KeyNorm] {
554            weights.insert(
555                TensorId::Vision {
556                    layer: layer_id,
557                    tensor,
558                },
559                ReferenceTensor::new(
560                    vec![layer.attention.head_dim as usize],
561                    vec![1.0; layer.attention.head_dim as usize],
562                )?,
563            );
564        }
565    }
566    let side = plan.pooling_kernel_size.max(1) as usize;
567    let output_tokens = output_tokens.unwrap_or(1) as usize;
568    let patch_count = side * side * output_tokens;
569    let mut patches = generated_tensor(&[patch_count, patch_width], 170, 0.5)?;
570    for value in &mut patches.data {
571        *value += 0.5;
572    }
573    let mut patch_positions = Vec::with_capacity(patch_count);
574    for y in 0..side {
575        for x in 0..side * output_tokens {
576            patch_positions.push([x as u32, y as u32]);
577        }
578    }
579    Ok(ReferenceVisionInput {
580        patches,
581        positions: patch_positions,
582        output_tokens,
583    })
584}
585
586fn add_hyper_head_fixture(
587    weights: &mut ReferenceWeights,
588    streams: usize,
589    hidden: usize,
590) -> Result<(), ReferenceError> {
591    if streams == 0 {
592        return Err(ReferenceError::InvalidPlan {
593            layer: None,
594            reason: "HyperConnections require at least one stream",
595        });
596    }
597    weights.insert(
598        TensorId::HyperHeadFunction,
599        generated_tensor(&[streams, streams * hidden], 90, 0.1)?,
600    );
601    weights.insert(
602        TensorId::HyperHeadBase,
603        generated_tensor(&[streams], 91, 0.05)?,
604    );
605    weights.insert(
606        TensorId::HyperHeadScale,
607        ReferenceTensor::new(vec![1], vec![0.2])?,
608    );
609    Ok(())
610}
611
612fn add_hyper_fixture(
613    weights: &mut ReferenceWeights,
614    layer: u32,
615    streams: usize,
616    hidden: usize,
617) -> Result<(), ReferenceError> {
618    if streams == 0 {
619        return Err(ReferenceError::InvalidPlan {
620            layer: Some(layer),
621            reason: "HyperConnections require at least one stream",
622        });
623    }
624    let rows = (2 + streams) * streams;
625    for (function, base, scale, salt) in [
626        (
627            LayerTensor::HyperAttentionFunction,
628            LayerTensor::HyperAttentionBase,
629            LayerTensor::HyperAttentionScale,
630            92,
631        ),
632        (
633            LayerTensor::HyperMlpFunction,
634            LayerTensor::HyperMlpBase,
635            LayerTensor::HyperMlpScale,
636            95,
637        ),
638    ] {
639        weights.insert(
640            layer_id(layer, function),
641            generated_tensor(&[rows, streams * hidden], salt + layer as u64 * 101, 0.1)?,
642        );
643        weights.insert(
644            layer_id(layer, base),
645            generated_tensor(&[rows], salt + 1 + layer as u64 * 101, 0.05)?,
646        );
647        weights.insert(
648            layer_id(layer, scale),
649            ReferenceTensor::new(vec![3], vec![0.2, 0.2, 0.2])?,
650        );
651    }
652    Ok(())
653}
654
655fn generated_tensor(
656    shape: &[usize],
657    salt: u64,
658    scale: f32,
659) -> Result<ReferenceTensor, ReferenceError> {
660    let elements = shape.iter().product();
661    let data = (0..elements)
662        .map(|index| {
663            let mut value = index as u64 ^ salt.wrapping_mul(0x9e37_79b9);
664            value ^= value >> 16;
665            value = value.wrapping_mul(0x45d9_f3b);
666            value ^= value >> 16;
667            let unit = (value as u32) as f32 / u32::MAX as f32;
668            (2.0 * unit - 1.0) * scale
669        })
670        .collect();
671    ReferenceTensor::new(shape.to_vec(), data)
672}
673
674fn add_full_attention_fixture(
675    weights: &mut ReferenceWeights,
676    layer: u32,
677    attention: &memra_gguf::model_plan::FullAttentionPlan,
678    hidden: usize,
679) -> Result<(), ReferenceError> {
680    let query_heads = attention.query_heads as usize;
681    let kv_heads = attention.kv_heads as usize;
682    let key_dim = attention.key_head_dim as usize;
683    let value_dim = attention.value_head_dim as usize;
684    let q_width = query_heads
685        * key_dim
686        * if attention.output_gate == AttentionGateKind::FusedQ {
687            2
688        } else {
689            1
690        };
691    for (tensor, output, input, salt) in [
692        (LayerTensor::Query, q_width, hidden, 10),
693        (LayerTensor::Key, kv_heads * key_dim, hidden, 11),
694        (
695            LayerTensor::AttentionOutput,
696            hidden,
697            query_heads * value_dim,
698            13,
699        ),
700    ] {
701        weights.insert(
702            layer_id(layer, tensor),
703            generated_tensor(
704                &[output, input],
705                salt + layer as u64 * 31,
706                1.0 / (input as f32).sqrt(),
707            )?,
708        );
709    }
710    if attention.value_projection == ValueProjection::Separate {
711        weights.insert(
712            layer_id(layer, LayerTensor::Value),
713            generated_tensor(
714                &[kv_heads * value_dim, hidden],
715                12 + layer as u64 * 31,
716                1.0 / (hidden as f32).sqrt(),
717            )?,
718        );
719    }
720    if attention.qk_norm != memra_gguf::model_plan::TensorPresence::Absent {
721        for tensor in [LayerTensor::QueryNorm, LayerTensor::KeyNorm] {
722            weights.insert(
723                layer_id(layer, tensor),
724                ReferenceTensor::new(vec![key_dim], vec![1.0; key_dim])?,
725            );
726        }
727    }
728    if attention.output_gate == AttentionGateKind::SeparateHead {
729        weights.insert(
730            layer_id(layer, LayerTensor::AttentionGate),
731            generated_tensor(
732                &[query_heads, hidden],
733                14 + layer as u64 * 31,
734                1.0 / (hidden as f32).sqrt(),
735            )?,
736        );
737    }
738    Ok(())
739}
740
741fn add_gdn_fixture(
742    weights: &mut ReferenceWeights,
743    layer: u32,
744    gdn: &memra_gguf::model_plan::GatedDeltaNetPlan,
745    hidden: usize,
746) -> Result<(), ReferenceError> {
747    let key_heads = gdn.key_heads as usize;
748    let value_heads = gdn.value_heads as usize;
749    let key_dim = gdn.key_head_dim as usize;
750    let value_dim = gdn.value_head_dim as usize;
751    let conv_width = 2 * key_heads * key_dim + value_heads * value_dim;
752    for (tensor, output, input, salt) in [
753        (LayerTensor::GdnQkv, conv_width, hidden, 40),
754        (LayerTensor::GdnGate, value_heads * value_dim, hidden, 41),
755        (LayerTensor::GdnBeta, value_heads, hidden, 42),
756        (LayerTensor::GdnAlpha, value_heads, hidden, 43),
757        (LayerTensor::GdnOutput, hidden, value_heads * value_dim, 44),
758    ] {
759        weights.insert(
760            layer_id(layer, tensor),
761            generated_tensor(
762                &[output, input],
763                salt + layer as u64 * 47,
764                1.0 / (input as f32).sqrt(),
765            )?,
766        );
767    }
768    weights.insert(
769        layer_id(layer, LayerTensor::GdnA),
770        ReferenceTensor::new(vec![value_heads], vec![-0.5; value_heads])?,
771    );
772    weights.insert(
773        layer_id(layer, LayerTensor::GdnDtBias),
774        ReferenceTensor::new(vec![value_heads], vec![0.0; value_heads])?,
775    );
776    weights.insert(
777        layer_id(layer, LayerTensor::GdnNorm),
778        ReferenceTensor::new(vec![value_dim], vec![1.0; value_dim])?,
779    );
780    weights.insert(
781        layer_id(layer, LayerTensor::GdnConv1d),
782        generated_tensor(
783            &[conv_width, gdn.conv_kernel as usize],
784            45 + layer as u64 * 47,
785            1.0 / (gdn.conv_kernel as f32).sqrt(),
786        )?,
787    );
788    Ok(())
789}
790
791fn add_mla_fixture(
792    weights: &mut ReferenceWeights,
793    layer: u32,
794    mla: &memra_gguf::model_plan::MlaAttentionPlan,
795    hidden: usize,
796) -> Result<(), ReferenceError> {
797    if let memra_gguf::model_plan::MlaAttentionPlan::CompressedKv { .. } = mla {
798        return add_compressed_mla_fixture(weights, layer, mla, hidden);
799    }
800    let memra_gguf::model_plan::MlaAttentionPlan::LatentKv {
801        query_heads,
802        q_lora_rank,
803        kv_lora_rank,
804        qk_head_dim,
805        rope_head_dim,
806        value_head_dim,
807        ..
808    } = mla.clone()
809    else {
810        return Err(ReferenceError::UnsupportedOperation {
811            layer: Some(layer),
812            operation: "compressed-KV MLA fixture",
813        });
814    };
815    let heads = query_heads as usize;
816    let q_rank = q_lora_rank as usize;
817    let kv_rank = kv_lora_rank as usize;
818    let qk_dim = qk_head_dim as usize;
819    let rope_dim = rope_head_dim as usize;
820    let nope_dim = qk_dim - rope_dim;
821    let value_dim = value_head_dim as usize;
822    for (tensor, shape, input, salt) in [
823        (LayerTensor::MlaQueryDown, vec![q_rank, hidden], hidden, 80),
824        (
825            LayerTensor::MlaQueryUp,
826            vec![heads * qk_dim, q_rank],
827            q_rank,
828            81,
829        ),
830        (
831            LayerTensor::MlaKvDown,
832            vec![kv_rank + rope_dim, hidden],
833            hidden,
834            82,
835        ),
836        (
837            LayerTensor::MlaKeyUp,
838            vec![heads, nope_dim, kv_rank],
839            kv_rank,
840            83,
841        ),
842        (
843            LayerTensor::MlaValueUp,
844            vec![heads, value_dim, kv_rank],
845            kv_rank,
846            84,
847        ),
848        (
849            LayerTensor::MlaOutput,
850            vec![hidden, heads * value_dim],
851            heads * value_dim,
852            85,
853        ),
854    ] {
855        weights.insert(
856            layer_id(layer, tensor),
857            generated_tensor(
858                &shape,
859                salt + layer as u64 * 71,
860                1.0 / (input as f32).sqrt(),
861            )?,
862        );
863    }
864    weights.insert(
865        layer_id(layer, LayerTensor::MlaQueryDownNorm),
866        ReferenceTensor::new(vec![q_rank], vec![1.0; q_rank])?,
867    );
868    weights.insert(
869        layer_id(layer, LayerTensor::MlaKvDownNorm),
870        ReferenceTensor::new(vec![kv_rank], vec![1.0; kv_rank])?,
871    );
872    Ok(())
873}
874
875fn add_compressed_mla_fixture(
876    weights: &mut ReferenceWeights,
877    layer: u32,
878    mla: &memra_gguf::model_plan::MlaAttentionPlan,
879    hidden: usize,
880) -> Result<(), ReferenceError> {
881    use memra_gguf::model_plan::{MlaAttentionPlan, SparseIndexPlan};
882
883    let MlaAttentionPlan::CompressedKv {
884        query_heads,
885        q_lora_rank,
886        latent_head_dim,
887        rope_head_dim,
888        output_lora_rank,
889        output_groups,
890        compressor,
891        sparse_index,
892        ..
893    } = mla
894    else {
895        unreachable!()
896    };
897    let heads = *query_heads as usize;
898    let q_rank = *q_lora_rank as usize;
899    let head_dim = *latent_head_dim as usize;
900    let rope_dim = *rope_head_dim as usize;
901    let output_rank = *output_lora_rank as usize;
902    let groups = *output_groups as usize;
903    if groups == 0 || heads % groups != 0 || rope_dim > head_dim {
904        return Err(ReferenceError::InvalidPlan {
905            layer: Some(layer),
906            reason: "compressed attention has invalid head or output-group geometry",
907        });
908    }
909    let group_width = heads / groups * head_dim;
910    for (tensor, shape, input, salt) in [
911        (LayerTensor::MlaQueryDown, vec![q_rank, hidden], hidden, 110),
912        (
913            LayerTensor::MlaQueryUp,
914            vec![heads * head_dim, q_rank],
915            q_rank,
916            111,
917        ),
918        (LayerTensor::MlaKvDown, vec![head_dim, hidden], hidden, 112),
919        (
920            LayerTensor::MlaOutputDown,
921            vec![groups * output_rank, group_width],
922            group_width,
923            113,
924        ),
925        (
926            LayerTensor::MlaOutput,
927            vec![hidden, groups * output_rank],
928            groups * output_rank,
929            114,
930        ),
931    ] {
932        weights.insert(
933            layer_id(layer, tensor),
934            generated_tensor(
935                &shape,
936                salt + layer as u64 * 131,
937                1.0 / (input as f32).sqrt(),
938            )?,
939        );
940    }
941    weights.insert(
942        layer_id(layer, LayerTensor::MlaQueryDownNorm),
943        ReferenceTensor::new(vec![q_rank], vec![1.0; q_rank])?,
944    );
945    weights.insert(
946        layer_id(layer, LayerTensor::MlaKvDownNorm),
947        ReferenceTensor::new(vec![head_dim], vec![1.0; head_dim])?,
948    );
949    weights.insert(
950        layer_id(layer, LayerTensor::AttentionSink),
951        generated_tensor(&[heads], 115 + layer as u64 * 131, 0.05)?,
952    );
953    if let Some(compressor) = compressor {
954        add_compressor_fixture(
955            weights,
956            layer,
957            hidden,
958            head_dim,
959            compressor.ratio as usize,
960            compressor.latent_dim as usize,
961            false,
962        )?;
963    }
964    match sparse_index {
965        SparseIndexPlan::None => {}
966        SparseIndexPlan::Own {
967            heads, head_dim, ..
968        } => {
969            let Some(compressor) = compressor else {
970                return Err(ReferenceError::InvalidPlan {
971                    layer: Some(layer),
972                    reason: "compressed sparse index requires a compressor ratio",
973                });
974            };
975            let index_heads = *heads as usize;
976            let index_dim = *head_dim as usize;
977            weights.insert(
978                layer_id(layer, LayerTensor::SparseQuery),
979                generated_tensor(
980                    &[index_heads * index_dim, q_rank],
981                    116 + layer as u64 * 131,
982                    1.0 / (q_rank as f32).sqrt(),
983                )?,
984            );
985            weights.insert(
986                layer_id(layer, LayerTensor::SparseProjection),
987                generated_tensor(
988                    &[index_heads, hidden],
989                    117 + layer as u64 * 131,
990                    1.0 / (hidden as f32).sqrt(),
991                )?,
992            );
993            add_compressor_fixture(
994                weights,
995                layer,
996                hidden,
997                index_dim,
998                compressor.ratio as usize,
999                2 * index_dim,
1000                true,
1001            )?;
1002        }
1003        SparseIndexPlan::SharedFromPrevious { .. } => {
1004            return Err(ReferenceError::UnsupportedOperation {
1005                layer: Some(layer),
1006                operation: "shared compressed sparse-index fixture",
1007            });
1008        }
1009    }
1010    Ok(())
1011}
1012
1013#[allow(clippy::too_many_arguments)]
1014fn add_compressor_fixture(
1015    weights: &mut ReferenceWeights,
1016    layer: u32,
1017    hidden: usize,
1018    output_dim: usize,
1019    ratio: usize,
1020    latent: usize,
1021    sparse: bool,
1022) -> Result<(), ReferenceError> {
1023    let (key_value, gate, norm, position, salt) = if sparse {
1024        (
1025            LayerTensor::SparseCompressorKeyValue,
1026            LayerTensor::SparseCompressorGate,
1027            LayerTensor::SparseCompressorNorm,
1028            LayerTensor::SparseCompressorPosition,
1029            121,
1030        )
1031    } else {
1032        (
1033            LayerTensor::KvCompressorKeyValue,
1034            LayerTensor::KvCompressorGate,
1035            LayerTensor::KvCompressorNorm,
1036            LayerTensor::KvCompressorPosition,
1037            118,
1038        )
1039    };
1040    for (tensor, offset) in [(key_value, 0), (gate, 1)] {
1041        weights.insert(
1042            layer_id(layer, tensor),
1043            generated_tensor(
1044                &[latent, hidden],
1045                salt + offset + layer as u64 * 131,
1046                1.0 / (hidden as f32).sqrt(),
1047            )?,
1048        );
1049    }
1050    weights.insert(
1051        layer_id(layer, norm),
1052        ReferenceTensor::new(vec![output_dim], vec![1.0; output_dim])?,
1053    );
1054    weights.insert(
1055        layer_id(layer, position),
1056        generated_tensor(&[ratio, latent], salt + 2 + layer as u64 * 131, 0.05)?,
1057    );
1058    Ok(())
1059}
1060
1061fn add_dense_mlp_fixture(
1062    weights: &mut ReferenceWeights,
1063    layer: u32,
1064    mlp: &memra_gguf::model_plan::DenseMlpPlan,
1065    hidden: usize,
1066) -> Result<(), ReferenceError> {
1067    let intermediate = mlp.intermediate_size as usize;
1068    for (tensor, output, input, salt) in [
1069        (LayerTensor::MlpGate, intermediate, hidden, 20),
1070        (LayerTensor::MlpUp, intermediate, hidden, 21),
1071        (LayerTensor::MlpDown, hidden, intermediate, 22),
1072    ] {
1073        weights.insert(
1074            layer_id(layer, tensor),
1075            generated_tensor(
1076                &[output, input],
1077                salt + layer as u64 * 31,
1078                1.0 / (input as f32).sqrt(),
1079            )?,
1080        );
1081    }
1082    Ok(())
1083}
1084
1085fn add_moe_fixture(
1086    weights: &mut ReferenceWeights,
1087    layer: u32,
1088    moe: &memra_gguf::model_plan::MoeMlpPlan,
1089    hidden: usize,
1090    vocab: usize,
1091) -> Result<(), ReferenceError> {
1092    let experts = moe.expert_count as usize;
1093    let selected = moe.experts_per_token as usize;
1094    let intermediate = moe.expert_intermediate_size as usize;
1095    if matches!(
1096        moe.router,
1097        memra_gguf::model_plan::RouterPlan::TokenIdHash { .. }
1098    ) {
1099        let mut table = Vec::with_capacity(vocab * selected);
1100        for token in 0..vocab {
1101            for rank in 0..selected {
1102                table.push(((token + rank) % experts) as f32);
1103            }
1104        }
1105        weights.insert(
1106            layer_id(layer, LayerTensor::MoeTokenToExpert),
1107            ReferenceTensor::new(vec![vocab, selected], table)?,
1108        );
1109    }
1110    weights.insert(
1111        layer_id(layer, LayerTensor::MoeRouter),
1112        generated_tensor(
1113            &[experts, hidden],
1114            60 + layer as u64 * 59,
1115            1.0 / (hidden as f32).sqrt(),
1116        )?,
1117    );
1118    if router_has_selection_bias(&moe.router) {
1119        weights.insert(
1120            layer_id(layer, LayerTensor::MoeRouterBias),
1121            generated_tensor(&[experts], 61 + layer as u64 * 59, 0.05)?,
1122        );
1123    }
1124    for (tensor, shape, input, salt) in [
1125        (
1126            LayerTensor::MoeExpertGateBank,
1127            vec![experts, intermediate, hidden],
1128            hidden,
1129            62,
1130        ),
1131        (
1132            LayerTensor::MoeExpertUpBank,
1133            vec![experts, intermediate, hidden],
1134            hidden,
1135            63,
1136        ),
1137        (
1138            LayerTensor::MoeExpertDownBank,
1139            vec![experts, hidden, intermediate],
1140            intermediate,
1141            64,
1142        ),
1143    ] {
1144        weights.insert(
1145            layer_id(layer, tensor),
1146            generated_tensor(
1147                &shape,
1148                salt + layer as u64 * 59,
1149                1.0 / (input as f32).sqrt(),
1150            )?,
1151        );
1152    }
1153    if let Some(shared) = moe.shared.as_ref() {
1154        let intermediate = shared.intermediate_size as usize;
1155        for (tensor, output, input, salt) in [
1156            (LayerTensor::SharedMlpGate, intermediate, hidden, 65),
1157            (LayerTensor::SharedMlpUp, intermediate, hidden, 66),
1158            (LayerTensor::SharedMlpDown, hidden, intermediate, 67),
1159        ] {
1160            weights.insert(
1161                layer_id(layer, tensor),
1162                generated_tensor(
1163                    &[output, input],
1164                    salt + layer as u64 * 59,
1165                    1.0 / (input as f32).sqrt(),
1166                )?,
1167            );
1168        }
1169        if shared.gated {
1170            weights.insert(
1171                layer_id(layer, LayerTensor::SharedMlpInputGate),
1172                generated_tensor(&[hidden], 68 + layer as u64 * 59, 0.2)?,
1173            );
1174        }
1175    }
1176    Ok(())
1177}
1178
1179fn add_gemma_parallel_moe_fixture(
1180    weights: &mut ReferenceWeights,
1181    layer: u32,
1182    moe: &memra_gguf::model_plan::MoeMlpPlan,
1183    hidden: usize,
1184) -> Result<(), ReferenceError> {
1185    let experts = moe.expert_count as usize;
1186    let intermediate = moe.expert_intermediate_size as usize;
1187    weights.insert(
1188        layer_id(layer, LayerTensor::MoeExpertGateUpBank),
1189        generated_tensor(
1190            &[experts, 2 * intermediate, hidden],
1191            180 + layer as u64 * 19,
1192            1.0 / (hidden as f32).sqrt(),
1193        )?,
1194    );
1195    weights.insert(
1196        layer_id(layer, LayerTensor::MoeRouterScale),
1197        ReferenceTensor::new(vec![hidden], vec![1.0; hidden])?,
1198    );
1199    weights.insert(
1200        layer_id(layer, LayerTensor::MoeExpertOutputScale),
1201        generated_tensor(&[experts], 181 + layer as u64 * 19, 0.2)?,
1202    );
1203    Ok(())
1204}
1205
1206pub fn execute(
1207    plan: &ModelPlan,
1208    weights: &ReferenceWeights,
1209    token_ids: &[u32],
1210) -> Result<ReferenceOutput, ReferenceError> {
1211    if token_ids.is_empty() {
1212        return Err(ReferenceError::EmptyInput);
1213    }
1214    let hidden = plan.hidden_size as usize;
1215    let vocab = plan.vocab_size as usize;
1216    let embedding = tensor(weights, &TensorId::TokenEmbedding, &[vocab, hidden])?;
1217    let embedded = embed_token_rows(plan, embedding, token_ids, vocab, hidden)?;
1218    execute_embedded(plan, weights, token_ids, embedding, embedded)
1219}
1220
1221pub fn execute_multimodal(
1222    plan: &ModelPlan,
1223    weights: &ReferenceWeights,
1224    token_ids: &[u32],
1225    vision_input: &ReferenceVisionInput,
1226) -> Result<ReferenceMultimodalOutput, ReferenceError> {
1227    if token_ids.is_empty() {
1228        return Err(ReferenceError::EmptyInput);
1229    }
1230    let injection = plan.multimodal.ok_or(ReferenceError::InvalidPlan {
1231        layer: None,
1232        reason: "multimodal input requires a vision-token injection plan",
1233    })?;
1234    let vision = execute_vision(plan, weights, vision_input)?;
1235    if vision.output_tokens != injection.tokens_per_image as usize {
1236        return Err(ReferenceError::InvalidPlan {
1237            layer: None,
1238            reason: "vision output token count does not match the injection plan",
1239        });
1240    }
1241    let placeholder_count = token_ids
1242        .iter()
1243        .filter(|&&token| token == injection.placeholder_token_id)
1244        .count();
1245    if placeholder_count != vision.output_tokens {
1246        return Err(ReferenceError::InvalidPlan {
1247            layer: None,
1248            reason: "image placeholder count does not match projected vision tokens",
1249        });
1250    }
1251    let hidden = plan.hidden_size as usize;
1252    let vocab = plan.vocab_size as usize;
1253    let embedding = tensor(weights, &TensorId::TokenEmbedding, &[vocab, hidden])?;
1254    let mut embedded = embed_token_rows(plan, embedding, token_ids, vocab, hidden)?;
1255    let mut vision_row = 0;
1256    for (position, &token) in token_ids.iter().enumerate() {
1257        if token == injection.placeholder_token_id {
1258            embedded[position * hidden..(position + 1) * hidden].copy_from_slice(
1259                &vision.projected_hidden[vision_row * hidden..(vision_row + 1) * hidden],
1260            );
1261            vision_row += 1;
1262        }
1263    }
1264    let language = execute_embedded(plan, weights, token_ids, embedding, embedded)?;
1265    Ok(ReferenceMultimodalOutput { language, vision })
1266}
1267
1268fn embed_token_rows(
1269    plan: &ModelPlan,
1270    embedding: &[f32],
1271    token_ids: &[u32],
1272    vocab: usize,
1273    hidden: usize,
1274) -> Result<Vec<f32>, ReferenceError> {
1275    let mut embedded = vec![0.0; token_ids.len() * hidden];
1276    for (position, &token) in token_ids.iter().enumerate() {
1277        let token = token as usize;
1278        if token >= vocab {
1279            return Err(ReferenceError::TokenOutOfRange {
1280                token: token as u32,
1281                vocab,
1282            });
1283        }
1284        embedded[position * hidden..(position + 1) * hidden]
1285            .copy_from_slice(&embedding[token * hidden..(token + 1) * hidden]);
1286        if plan.embedding_scale != 1.0 {
1287            for value in &mut embedded[position * hidden..(position + 1) * hidden] {
1288                *value *= plan.embedding_scale;
1289            }
1290        }
1291    }
1292    Ok(embedded)
1293}
1294
1295fn execute_embedded(
1296    plan: &ModelPlan,
1297    weights: &ReferenceWeights,
1298    token_ids: &[u32],
1299    embedding: &[f32],
1300    embedded: Vec<f32>,
1301) -> Result<ReferenceOutput, ReferenceError> {
1302    let tokens = token_ids.len();
1303    let hidden = plan.hidden_size as usize;
1304    let vocab = plan.vocab_size as usize;
1305    if embedded.len() != tokens * hidden {
1306        return Err(ReferenceError::InvalidPlan {
1307            layer: None,
1308            reason: "embedded language input does not match tokens x hidden",
1309        });
1310    }
1311    let hyper = hyper_topology(plan)?;
1312    let mut x = hyper.map_or(embedded.clone(), |(streams, _, _)| {
1313        memra_gguf::dsv4_forward::hc_expand(&embedded, tokens, streams, hidden)
1314    });
1315
1316    let mut state = Vec::with_capacity(plan.layers.len());
1317    let dspark = plan.drafter.as_ref().map(|drafter| match drafter {
1318        memra_gguf::model_plan::DrafterPlan::Dspark(plan) => plan,
1319    });
1320    let mut draft_taps = dspark.map(|plan| vec![None; plan.target_layer_ids.len()]);
1321    for layer in &plan.layers {
1322        let (next, layer_state) =
1323            execute_layer(layer, weights, &x, token_ids, tokens, hidden, vocab)?;
1324        x = next;
1325        if let (Some(dspark), Some(taps)) = (dspark, draft_taps.as_mut()) {
1326            if let Some(target) = dspark
1327                .target_layer_ids
1328                .iter()
1329                .position(|&target| target == layer.index)
1330            {
1331                taps[target] = Some(collapse_stream_mean(&x, tokens, hidden, hyper)?);
1332            }
1333        }
1334        state.push(layer_state);
1335    }
1336    let trunk_hidden = x.clone();
1337    let x = if let Some((streams, epsilon, _)) = hyper {
1338        collapse_hyper_head(weights, &x, tokens, streams, hidden, plan, epsilon)?
1339    } else {
1340        x
1341    };
1342    let x = rms_norm(
1343        &x,
1344        tokens,
1345        hidden,
1346        tensor(weights, &TensorId::OutputNorm, &[hidden])?,
1347        plan.output_norm.epsilon,
1348    );
1349    let output = weights
1350        .get(&TensorId::OutputProjection)
1351        .map(|tensor| tensor_checked(&TensorId::OutputProjection, tensor, &[vocab, hidden]))
1352        .transpose()?
1353        .unwrap_or(embedding);
1354    let mut logits = linear(&x, output, tokens, hidden, vocab);
1355    apply_logits_transforms(&mut logits, vocab, &plan.logits);
1356    let draft = match (dspark, draft_taps) {
1357        (Some(dspark), Some(taps)) => Some(execute_dspark(
1358            dspark,
1359            weights,
1360            token_ids,
1361            embedding,
1362            output,
1363            &plan.logits,
1364            plan.output_norm.epsilon,
1365            hidden,
1366            vocab,
1367            taps,
1368        )?),
1369        _ => None,
1370    };
1371    let mtp = execute_mtp(
1372        plan,
1373        weights,
1374        token_ids,
1375        embedding,
1376        &trunk_hidden,
1377        tokens,
1378        hidden,
1379        vocab,
1380        output,
1381    )?;
1382    Ok(ReferenceOutput {
1383        logits,
1384        tokens,
1385        vocab,
1386        state: ReferenceState { layers: state },
1387        mtp,
1388        draft,
1389    })
1390}
1391
1392pub fn execute_vision(
1393    plan: &ModelPlan,
1394    weights: &ReferenceWeights,
1395    input: &ReferenceVisionInput,
1396) -> Result<ReferenceVisionOutput, ReferenceError> {
1397    let Some(vision) = plan.vision.as_ref() else {
1398        return Err(ReferenceError::InvalidPlan {
1399            layer: None,
1400            reason: "vision input requires a vision subplan",
1401        });
1402    };
1403    if vision.clipped_linears {
1404        return Err(ReferenceError::UnsupportedOperation {
1405            layer: None,
1406            operation: "clipped vision linears",
1407        });
1408    }
1409    let patches = input.positions.len();
1410    let hidden = vision.hidden_size as usize;
1411    let patch_width =
1412        (vision.patch.channels * vision.patch.patch_size * vision.patch.patch_size) as usize;
1413    if input.patches.shape != [patches, patch_width]
1414        || input.output_tokens == 0
1415        || input.output_tokens > patches
1416    {
1417        return Err(ReferenceError::InvalidPlan {
1418            layer: None,
1419            reason: "vision patch input shape or output-token count is invalid",
1420        });
1421    }
1422    let mut normalized_patches = input.patches.data.clone();
1423    for value in &mut normalized_patches {
1424        *value = 2.0 * (*value - 0.5);
1425    }
1426    let mut x = linear(
1427        &normalized_patches,
1428        tensor(
1429            weights,
1430            &TensorId::Vision {
1431                layer: None,
1432                tensor: VisionTensor::PatchProjection,
1433            },
1434            &[hidden, patch_width],
1435        )?,
1436        patches,
1437        patch_width,
1438        hidden,
1439    );
1440    let position_table = tensor(
1441        weights,
1442        &TensorId::Vision {
1443            layer: None,
1444            tensor: VisionTensor::PositionEmbedding,
1445        },
1446        &[
1447            vision.patch.position_axes as usize,
1448            vision.patch.position_embedding_size as usize,
1449            hidden,
1450        ],
1451    )?;
1452    for (patch, position) in input.positions.iter().enumerate() {
1453        for (axis, &coordinate) in position.iter().enumerate() {
1454            let coordinate = coordinate as usize;
1455            if axis >= vision.patch.position_axes as usize
1456                || coordinate >= vision.patch.position_embedding_size as usize
1457            {
1458                return Err(ReferenceError::InvalidPlan {
1459                    layer: None,
1460                    reason: "vision patch position is outside the embedding table",
1461                });
1462            }
1463            let source =
1464                (axis * vision.patch.position_embedding_size as usize + coordinate) * hidden;
1465            for column in 0..hidden {
1466                x[patch * hidden + column] += position_table[source + column];
1467            }
1468        }
1469    }
1470    for layer in &vision.layers {
1471        x = execute_vision_layer(layer, weights, &x, &input.positions, patches, hidden)?;
1472    }
1473    let encoder_hidden = x.clone();
1474    let pooled_hidden = vision_pool(&x, &input.positions, patches, input.output_tokens, hidden)?;
1475    let mut standardized = pooled_hidden.clone();
1476    if vision.standardize {
1477        let bias = tensor(
1478            weights,
1479            &TensorId::Vision {
1480                layer: None,
1481                tensor: VisionTensor::StandardizeBias,
1482            },
1483            &[hidden],
1484        )?;
1485        let scale = tensor(
1486            weights,
1487            &TensorId::Vision {
1488                layer: None,
1489                tensor: VisionTensor::StandardizeScale,
1490            },
1491            &[hidden],
1492        )?;
1493        for row in standardized.chunks_exact_mut(hidden) {
1494            for column in 0..hidden {
1495                row[column] = (row[column] - bias[column]) * scale[column];
1496            }
1497        }
1498    }
1499    let standardized = rms_norm(
1500        &standardized,
1501        input.output_tokens,
1502        hidden,
1503        &vec![1.0; hidden],
1504        vision.layers[0].input_norm.epsilon,
1505    );
1506    let projection_size = vision.projection_output_size as usize;
1507    let projected_hidden = linear(
1508        &standardized,
1509        tensor(
1510            weights,
1511            &TensorId::Vision {
1512                layer: None,
1513                tensor: VisionTensor::OutputProjection,
1514            },
1515            &[projection_size, hidden],
1516        )?,
1517        input.output_tokens,
1518        hidden,
1519        projection_size,
1520    );
1521    Ok(ReferenceVisionOutput {
1522        encoder_hidden,
1523        pooled_hidden,
1524        projected_hidden,
1525        patch_count: patches,
1526        output_tokens: input.output_tokens,
1527        hidden_size: hidden,
1528        projection_size,
1529    })
1530}
1531
1532fn execute_vision_layer(
1533    plan: &memra_gguf::model_plan::VisionLayerPlan,
1534    weights: &ReferenceWeights,
1535    input: &[f32],
1536    positions: &[[u32; 2]],
1537    tokens: usize,
1538    hidden: usize,
1539) -> Result<Vec<f32>, ReferenceError> {
1540    let id = |tensor| TensorId::Vision {
1541        layer: Some(plan.index),
1542        tensor,
1543    };
1544    let attention_input = rms_norm(
1545        input,
1546        tokens,
1547        hidden,
1548        tensor(weights, &id(VisionTensor::InputNorm), &[hidden])?,
1549        plan.input_norm.epsilon,
1550    );
1551    let query_heads = plan.attention.query_heads as usize;
1552    let kv_heads = plan.attention.kv_heads as usize;
1553    let head_dim = plan.attention.head_dim as usize;
1554    if query_heads == 0 || kv_heads == 0 || query_heads % kv_heads != 0 {
1555        return Err(ReferenceError::InvalidPlan {
1556            layer: Some(plan.index),
1557            reason: "vision attention has invalid query/KV head grouping",
1558        });
1559    }
1560    let mut query = linear(
1561        &attention_input,
1562        tensor(
1563            weights,
1564            &id(VisionTensor::Query),
1565            &[query_heads * head_dim, hidden],
1566        )?,
1567        tokens,
1568        hidden,
1569        query_heads * head_dim,
1570    );
1571    let mut key = linear(
1572        &attention_input,
1573        tensor(
1574            weights,
1575            &id(VisionTensor::Key),
1576            &[kv_heads * head_dim, hidden],
1577        )?,
1578        tokens,
1579        hidden,
1580        kv_heads * head_dim,
1581    );
1582    let mut value = linear(
1583        &attention_input,
1584        tensor(
1585            weights,
1586            &id(VisionTensor::Value),
1587            &[kv_heads * head_dim, hidden],
1588        )?,
1589        tokens,
1590        hidden,
1591        kv_heads * head_dim,
1592    );
1593    apply_optional_head_norm(
1594        weights,
1595        id(VisionTensor::QueryNorm),
1596        &mut query,
1597        tokens * query_heads,
1598        head_dim,
1599        memra_gguf::model_plan::TensorPresence::Required,
1600        plan.input_norm.epsilon,
1601    )?;
1602    apply_optional_head_norm(
1603        weights,
1604        id(VisionTensor::KeyNorm),
1605        &mut key,
1606        tokens * kv_heads,
1607        head_dim,
1608        memra_gguf::model_plan::TensorPresence::Required,
1609        plan.input_norm.epsilon,
1610    )?;
1611    value = rms_norm(
1612        &value,
1613        tokens * kv_heads,
1614        head_dim,
1615        &vec![1.0; head_dim],
1616        plan.input_norm.epsilon,
1617    );
1618    apply_vision_rope(
1619        &mut query,
1620        tokens,
1621        query_heads,
1622        head_dim,
1623        positions,
1624        plan.attention.rope.base,
1625    )?;
1626    apply_vision_rope(
1627        &mut key,
1628        tokens,
1629        kv_heads,
1630        head_dim,
1631        positions,
1632        plan.attention.rope.base,
1633    )?;
1634    let repeat = query_heads / kv_heads;
1635    let mut attended = vec![0.0; tokens * query_heads * head_dim];
1636    for token in 0..tokens {
1637        for head in 0..query_heads {
1638            let kv_head = head / repeat;
1639            let mut scores = Vec::with_capacity(tokens);
1640            for source in 0..tokens {
1641                let mut score = 0.0;
1642                for column in 0..head_dim {
1643                    score += query[(token * query_heads + head) * head_dim + column]
1644                        * key[(source * kv_heads + kv_head) * head_dim + column];
1645                }
1646                scores.push(score);
1647            }
1648            softmax_in_place(&mut scores);
1649            for (source, probability) in scores.into_iter().enumerate() {
1650                for column in 0..head_dim {
1651                    attended[(token * query_heads + head) * head_dim + column] +=
1652                        probability * value[(source * kv_heads + kv_head) * head_dim + column];
1653                }
1654            }
1655        }
1656    }
1657    let attention = linear(
1658        &attended,
1659        tensor(
1660            weights,
1661            &id(VisionTensor::AttentionOutput),
1662            &[hidden, query_heads * head_dim],
1663        )?,
1664        tokens,
1665        query_heads * head_dim,
1666        hidden,
1667    );
1668    let attention = rms_norm(
1669        &attention,
1670        tokens,
1671        hidden,
1672        tensor(weights, &id(VisionTensor::PostAttentionNorm), &[hidden])?,
1673        plan.post_attention_norm.epsilon,
1674    );
1675    let mut residual = input.to_vec();
1676    add_in_place(&mut residual, &attention);
1677    let mlp_input = rms_norm(
1678        &residual,
1679        tokens,
1680        hidden,
1681        tensor(weights, &id(VisionTensor::PreMlpNorm), &[hidden])?,
1682        plan.pre_mlp_norm.epsilon,
1683    );
1684    let intermediate = plan.mlp.intermediate_size as usize;
1685    let gate = linear(
1686        &mlp_input,
1687        tensor(weights, &id(VisionTensor::MlpGate), &[intermediate, hidden])?,
1688        tokens,
1689        hidden,
1690        intermediate,
1691    );
1692    let up = linear(
1693        &mlp_input,
1694        tensor(weights, &id(VisionTensor::MlpUp), &[intermediate, hidden])?,
1695        tokens,
1696        hidden,
1697        intermediate,
1698    );
1699    let mut activated = vec![0.0; gate.len()];
1700    for index in 0..activated.len() {
1701        activated[index] = activate_pair(&plan.mlp.activation, gate[index], up[index], plan.index)?;
1702    }
1703    let mlp = linear(
1704        &activated,
1705        tensor(weights, &id(VisionTensor::MlpDown), &[hidden, intermediate])?,
1706        tokens,
1707        intermediate,
1708        hidden,
1709    );
1710    let mlp = rms_norm(
1711        &mlp,
1712        tokens,
1713        hidden,
1714        tensor(weights, &id(VisionTensor::PostMlpNorm), &[hidden])?,
1715        plan.post_mlp_norm.epsilon,
1716    );
1717    add_in_place(&mut residual, &mlp);
1718    Ok(residual)
1719}
1720
1721fn apply_vision_rope(
1722    values: &mut [f32],
1723    tokens: usize,
1724    heads: usize,
1725    head_dim: usize,
1726    positions: &[[u32; 2]],
1727    base: f32,
1728) -> Result<(), ReferenceError> {
1729    let axes = 2;
1730    let chunk = head_dim / axes;
1731    if head_dim % axes != 0 || chunk % 2 != 0 || positions.len() != tokens {
1732        return Err(ReferenceError::InvalidPlan {
1733            layer: None,
1734            reason: "vision 2D RoPE requires even per-axis head chunks",
1735        });
1736    }
1737    let half = chunk / 2;
1738    for token in 0..tokens {
1739        for head in 0..heads {
1740            let row = (token * heads + head) * head_dim;
1741            for axis in 0..axes {
1742                let start = row + axis * chunk;
1743                let position = positions[token][axis] as f32;
1744                for pair in 0..half {
1745                    let angle = position / base.powf((2 * pair) as f32 / chunk as f32);
1746                    let (sin, cos) = angle.sin_cos();
1747                    let left = values[start + pair];
1748                    let right = values[start + half + pair];
1749                    values[start + pair] = left * cos - right * sin;
1750                    values[start + half + pair] = left * sin + right * cos;
1751                }
1752            }
1753        }
1754    }
1755    Ok(())
1756}
1757
1758fn vision_pool(
1759    hidden_states: &[f32],
1760    positions: &[[u32; 2]],
1761    patches: usize,
1762    output_tokens: usize,
1763    hidden: usize,
1764) -> Result<Vec<f32>, ReferenceError> {
1765    if patches % output_tokens != 0 {
1766        return Err(ReferenceError::InvalidPlan {
1767            layer: None,
1768            reason: "vision pooling ratio must divide the patch count",
1769        });
1770    }
1771    let area = patches / output_tokens;
1772    let kernel = (area as f32).sqrt() as usize;
1773    if kernel * kernel != area {
1774        return Err(ReferenceError::InvalidPlan {
1775            layer: None,
1776            reason: "vision pooling ratio must be a square kernel",
1777        });
1778    }
1779    let max_x = positions
1780        .iter()
1781        .map(|position| position[0] as usize)
1782        .max()
1783        .unwrap_or(0)
1784        + 1;
1785    let grid_width = max_x / kernel;
1786    let mut output = vec![0.0; output_tokens * hidden];
1787    for patch in 0..patches {
1788        let target = positions[patch][0] as usize / kernel
1789            + grid_width * (positions[patch][1] as usize / kernel);
1790        if target >= output_tokens {
1791            return Err(ReferenceError::InvalidPlan {
1792                layer: None,
1793                reason: "vision patch positions do not fit the pooled grid",
1794            });
1795        }
1796        for column in 0..hidden {
1797            output[target * hidden + column] +=
1798                hidden_states[patch * hidden + column] / area as f32;
1799        }
1800    }
1801    let scale = (hidden as f32).sqrt();
1802    for value in &mut output {
1803        *value *= scale;
1804    }
1805    Ok(output)
1806}
1807
1808fn collapse_stream_mean(
1809    x: &[f32],
1810    tokens: usize,
1811    hidden: usize,
1812    hyper: Option<(usize, f32, u32)>,
1813) -> Result<Vec<f32>, ReferenceError> {
1814    let Some((streams, _, _)) = hyper else {
1815        if x.len() != tokens * hidden {
1816            return Err(ReferenceError::InvalidPlan {
1817                layer: None,
1818                reason: "single-stream DSpark tap has invalid shape",
1819            });
1820        }
1821        return Ok(x.to_vec());
1822    };
1823    if x.len() != tokens * streams * hidden {
1824        return Err(ReferenceError::InvalidPlan {
1825            layer: None,
1826            reason: "HyperConnections DSpark tap has invalid shape",
1827        });
1828    }
1829    let mut output = vec![0.0; tokens * hidden];
1830    for token in 0..tokens {
1831        for stream in 0..streams {
1832            for column in 0..hidden {
1833                output[token * hidden + column] +=
1834                    x[(token * streams + stream) * hidden + column] / streams as f32;
1835            }
1836        }
1837    }
1838    Ok(output)
1839}
1840
1841#[allow(clippy::too_many_arguments)]
1842fn execute_dspark(
1843    plan: &memra_gguf::model_plan::DsparkPlan,
1844    weights: &ReferenceWeights,
1845    token_ids: &[u32],
1846    embedding: &[f32],
1847    output_projection: &[f32],
1848    logits_transforms: &[LogitsTransform],
1849    norm_epsilon: f32,
1850    hidden: usize,
1851    vocab: usize,
1852    taps: Vec<Option<Vec<f32>>>,
1853) -> Result<ReferenceDraftOutput, ReferenceError> {
1854    use memra_gguf::dsv4_forward::{hc_expand, hc_head, matmul, rmsnorm};
1855
1856    let tokens = token_ids.len();
1857    let block_size = plan.block_size as usize;
1858    let rank = plan.markov_rank as usize;
1859    if tokens < 2
1860        || block_size == 0
1861        || plan.blocks.is_empty()
1862        || taps.len() != plan.target_layer_ids.len()
1863        || plan.noise_token_id as usize >= vocab
1864    {
1865        return Err(ReferenceError::InvalidPlan {
1866            layer: None,
1867            reason: "DSpark execution requires a primed prompt and valid drafter geometry",
1868        });
1869    }
1870    let streams = match plan.blocks[0].residual {
1871        ResidualTopology::HyperConnections { streams, .. } if streams > 0 => streams as usize,
1872        _ => {
1873            return Err(ReferenceError::InvalidPlan {
1874                layer: Some(plan.blocks[0].index),
1875                reason: "DSpark blocks require HyperConnections",
1876            });
1877        }
1878    };
1879    let mut main_hidden = vec![0.0; tokens * taps.len() * hidden];
1880    for (target, tap) in taps.into_iter().enumerate() {
1881        let Some(tap) = tap else {
1882            return Err(ReferenceError::InvalidPlan {
1883                layer: None,
1884                reason: "DSpark target layer was not captured from the trunk",
1885            });
1886        };
1887        if tap.len() != tokens * hidden {
1888            return Err(ReferenceError::InvalidPlan {
1889                layer: None,
1890                reason: "DSpark trunk tap has invalid shape",
1891            });
1892        }
1893        for token in 0..tokens {
1894            main_hidden[(token * plan.target_layer_ids.len() + target) * hidden
1895                ..(token * plan.target_layer_ids.len() + target + 1) * hidden]
1896                .copy_from_slice(&tap[token * hidden..(token + 1) * hidden]);
1897        }
1898    }
1899    let main_x = rmsnorm(
1900        &matmul(
1901            &main_hidden,
1902            tokens,
1903            plan.target_layer_ids.len() * hidden,
1904            tensor(
1905                weights,
1906                &TensorId::Dspark(DsparkTensor::MainProjection),
1907                &[hidden, plan.target_layer_ids.len() * hidden],
1908            )?,
1909            hidden,
1910        ),
1911        tensor(
1912            weights,
1913            &TensorId::Dspark(DsparkTensor::MainNorm),
1914            &[hidden],
1915        )?,
1916        norm_epsilon,
1917    );
1918    let rings = plan
1919        .blocks
1920        .iter()
1921        .map(|block| dspark_prime_ring(block, weights, &main_x, tokens, hidden, norm_epsilon))
1922        .collect::<Result<Vec<_>, _>>()?;
1923
1924    let input_token = *token_ids.last().unwrap();
1925    let mut draft_ids = vec![plan.noise_token_id; block_size];
1926    draft_ids[0] = input_token;
1927    let mut embedded = vec![0.0; block_size * hidden];
1928    for (position, &token) in draft_ids.iter().enumerate() {
1929        let token = token as usize;
1930        embedded[position * hidden..(position + 1) * hidden]
1931            .copy_from_slice(&embedding[token * hidden..(token + 1) * hidden]);
1932    }
1933    let mut draft_hidden = hc_expand(&embedded, block_size, streams, hidden);
1934    for (block, ring) in plan.blocks.iter().zip(&rings) {
1935        draft_hidden = execute_dspark_layer(
1936            block,
1937            weights,
1938            &draft_hidden,
1939            ring,
1940            tokens - 1,
1941            block_size,
1942            hidden,
1943            vocab,
1944        )?;
1945    }
1946    let head_set = memra_gguf::dsv4_forward::HcSet {
1947        rows: streams,
1948        fn_w: tensor(
1949            weights,
1950            &TensorId::Dspark(DsparkTensor::HeadHyperFunction),
1951            &[streams, streams * hidden],
1952        )?
1953        .to_vec(),
1954        base: tensor(
1955            weights,
1956            &TensorId::Dspark(DsparkTensor::HeadHyperBase),
1957            &[streams],
1958        )?
1959        .to_vec(),
1960        scale: tensor(
1961            weights,
1962            &TensorId::Dspark(DsparkTensor::HeadHyperScale),
1963            &[1],
1964        )?
1965        .to_vec(),
1966    };
1967    let hc_epsilon = match plan.blocks[0].residual {
1968        ResidualTopology::HyperConnections { epsilon, .. } => epsilon,
1969        _ => unreachable!(),
1970    };
1971    let collapsed = hc_head(
1972        &draft_hidden,
1973        block_size,
1974        streams,
1975        hidden,
1976        &head_set,
1977        norm_epsilon,
1978        hc_epsilon,
1979    );
1980    let normalized = rmsnorm(
1981        &collapsed,
1982        tensor(
1983            weights,
1984            &TensorId::Dspark(DsparkTensor::OutputNorm),
1985            &[hidden],
1986        )?,
1987        norm_epsilon,
1988    );
1989    let mut logits = matmul(&normalized, block_size, hidden, output_projection, vocab);
1990    apply_logits_transforms(&mut logits, vocab, logits_transforms);
1991
1992    let markov_embedding = tensor(
1993        weights,
1994        &TensorId::Dspark(DsparkTensor::MarkovEmbedding),
1995        &[vocab, rank],
1996    )?;
1997    let markov_output = tensor(
1998        weights,
1999        &TensorId::Dspark(DsparkTensor::MarkovOutput),
2000        &[vocab, rank],
2001    )?;
2002    let confidence_weight = tensor(
2003        weights,
2004        &TensorId::Dspark(DsparkTensor::ConfidenceProjection),
2005        &[1, hidden + rank],
2006    )?;
2007    let mut output_ids = vec![input_token];
2008    let mut confidence = Vec::with_capacity(block_size);
2009    for position in 0..block_size {
2010        let previous = output_ids[position] as usize;
2011        let markov = &markov_embedding[previous * rank..(previous + 1) * rank];
2012        let row = &mut logits[position * vocab..(position + 1) * vocab];
2013        for token in 0..vocab {
2014            row[token] += memra_gguf::dsv4_forward::dot(
2015                markov,
2016                &markov_output[token * rank..(token + 1) * rank],
2017            );
2018        }
2019        let next = row
2020            .iter()
2021            .enumerate()
2022            .max_by(|(left_index, left), (right_index, right)| {
2023                left.total_cmp(right)
2024                    .then_with(|| right_index.cmp(left_index))
2025            })
2026            .map(|(index, _)| index as u32)
2027            .unwrap();
2028        output_ids.push(next);
2029        let mut confidence_input = Vec::with_capacity(hidden + rank);
2030        confidence_input.extend_from_slice(&collapsed[position * hidden..(position + 1) * hidden]);
2031        confidence_input.extend_from_slice(markov);
2032        confidence.push(memra_gguf::dsv4_forward::dot(
2033            &confidence_input,
2034            confidence_weight,
2035        ));
2036    }
2037    Ok(ReferenceDraftOutput {
2038        input_token,
2039        output_ids,
2040        confidence,
2041        logits,
2042        hidden: collapsed,
2043        block_size,
2044    })
2045}
2046
2047fn dspark_prime_ring(
2048    layer: &memra_gguf::model_plan::LayerPlan,
2049    weights: &ReferenceWeights,
2050    main_x: &[f32],
2051    tokens: usize,
2052    hidden: usize,
2053    epsilon: f32,
2054) -> Result<Vec<f32>, ReferenceError> {
2055    use memra_gguf::dsv4_forward::{ActQuantVariant, apply_rope, matmul, rmsnorm};
2056    use memra_gguf::model_plan::{MlaAttentionPlan, RopeFactors, SparseIndexPlan};
2057
2058    let AttentionPlan::Mla(MlaAttentionPlan::CompressedKv {
2059        latent_head_dim,
2060        rope_head_dim,
2061        window,
2062        rope,
2063        compressor: None,
2064        sparse_index: SparseIndexPlan::None,
2065        ..
2066    }) = &layer.attention
2067    else {
2068        return Err(ReferenceError::InvalidPlan {
2069            layer: Some(layer.index),
2070            reason: "DSpark blocks require uncompressed window-only attention",
2071        });
2072    };
2073    if !matches!(rope.factors, RopeFactors::None) {
2074        return Err(ReferenceError::InvalidPlan {
2075            layer: Some(layer.index),
2076            reason: "DSpark block RoPE must not use scaling factors",
2077        });
2078    }
2079    let head_dim = *latent_head_dim as usize;
2080    let rope_dim = *rope_head_dim as usize;
2081    if head_dim <= rope_dim || (head_dim - rope_dim) % 64 != 0 {
2082        return Err(ReferenceError::InvalidPlan {
2083            layer: Some(layer.index),
2084            reason: "DSpark block has invalid KV quantization geometry",
2085        });
2086    }
2087    let frequencies = memra_gguf::dsv4_forward::precompute_freqs_cis(
2088        rope_dim,
2089        tokens + 1,
2090        0,
2091        rope.base,
2092        1.0,
2093        32.0,
2094        1.0,
2095    );
2096    let mut key_value = rmsnorm(
2097        &matmul(
2098            main_x,
2099            tokens,
2100            hidden,
2101            tensor(
2102                weights,
2103                &layer_id(layer.index, LayerTensor::MlaKvDown),
2104                &[head_dim, hidden],
2105            )?,
2106            head_dim,
2107        ),
2108        tensor(
2109            weights,
2110            &layer_id(layer.index, LayerTensor::MlaKvDownNorm),
2111            &[head_dim],
2112        )?,
2113        epsilon,
2114    );
2115    let positions: Vec<_> = (0..tokens).collect();
2116    apply_rope(
2117        &mut key_value,
2118        tokens,
2119        1,
2120        head_dim,
2121        rope_dim,
2122        &frequencies,
2123        &positions,
2124        false,
2125    );
2126    for row in key_value.chunks_exact_mut(head_dim) {
2127        memra_gguf::dsv4_forward::act_quant(
2128            &mut row[..head_dim - rope_dim],
2129            64,
2130            ActQuantVariant::RefFp8Round,
2131        );
2132    }
2133    let window = *window as usize;
2134    let mut ring = vec![0.0; window * head_dim];
2135    for position in tokens.saturating_sub(window)..tokens {
2136        ring[(position % window) * head_dim..(position % window + 1) * head_dim]
2137            .copy_from_slice(&key_value[position * head_dim..(position + 1) * head_dim]);
2138    }
2139    Ok(ring)
2140}
2141
2142#[allow(clippy::too_many_arguments)]
2143fn execute_dspark_layer(
2144    layer: &memra_gguf::model_plan::LayerPlan,
2145    weights: &ReferenceWeights,
2146    input: &[f32],
2147    ring: &[f32],
2148    start_position: usize,
2149    block_size: usize,
2150    hidden: usize,
2151    vocab: usize,
2152) -> Result<Vec<f32>, ReferenceError> {
2153    let ResidualTopology::HyperConnections {
2154        streams,
2155        epsilon,
2156        sinkhorn_iterations,
2157    } = layer.residual
2158    else {
2159        return Err(ReferenceError::InvalidPlan {
2160            layer: Some(layer.index),
2161            reason: "DSpark block requires HyperConnections",
2162        });
2163    };
2164    let streams = streams as usize;
2165    let attention_set = hyper_set(
2166        weights,
2167        layer.index,
2168        streams,
2169        hidden,
2170        LayerTensor::HyperAttentionFunction,
2171        LayerTensor::HyperAttentionBase,
2172        LayerTensor::HyperAttentionScale,
2173    )?;
2174    let (attention_input, post, combination) = memra_gguf::dsv4_forward::hc_pre(
2175        input,
2176        block_size,
2177        streams,
2178        hidden,
2179        &attention_set,
2180        sinkhorn_iterations,
2181        epsilon,
2182    );
2183    let attention_input = rms_norm(
2184        &attention_input,
2185        block_size,
2186        hidden,
2187        tensor(
2188            weights,
2189            &layer_id(layer.index, LayerTensor::PreAttentionNorm),
2190            &[hidden],
2191        )?,
2192        layer.pre_attention_norm.epsilon,
2193    );
2194    let attention = dspark_attention(
2195        layer,
2196        weights,
2197        &attention_input,
2198        ring,
2199        start_position,
2200        block_size,
2201        hidden,
2202    )?;
2203    let attention_residual = memra_gguf::dsv4_forward::hc_post(
2204        &attention,
2205        input,
2206        block_size,
2207        streams,
2208        hidden,
2209        &post,
2210        &combination,
2211    );
2212    let mlp_set = hyper_set(
2213        weights,
2214        layer.index,
2215        streams,
2216        hidden,
2217        LayerTensor::HyperMlpFunction,
2218        LayerTensor::HyperMlpBase,
2219        LayerTensor::HyperMlpScale,
2220    )?;
2221    let (mlp_input, post, combination) = memra_gguf::dsv4_forward::hc_pre(
2222        &attention_residual,
2223        block_size,
2224        streams,
2225        hidden,
2226        &mlp_set,
2227        sinkhorn_iterations,
2228        epsilon,
2229    );
2230    let mlp_input = rms_norm(
2231        &mlp_input,
2232        block_size,
2233        hidden,
2234        tensor(
2235            weights,
2236            &layer_id(layer.index, LayerTensor::PreMlpNorm),
2237            &[hidden],
2238        )?,
2239        layer.pre_mlp_norm.epsilon,
2240    );
2241    let zeros = vec![0; block_size];
2242    let mlp = match &layer.mlp {
2243        MlpPlan::Dense(mlp) => {
2244            dense_mlp(layer.index, mlp, weights, &mlp_input, block_size, hidden)?
2245        }
2246        MlpPlan::Moe(moe) => moe_mlp(
2247            layer.index,
2248            moe,
2249            weights,
2250            &mlp_input,
2251            &zeros,
2252            block_size,
2253            hidden,
2254            vocab,
2255        )?,
2256    };
2257    Ok(memra_gguf::dsv4_forward::hc_post(
2258        &mlp,
2259        &attention_residual,
2260        block_size,
2261        streams,
2262        hidden,
2263        &post,
2264        &combination,
2265    ))
2266}
2267
2268#[allow(clippy::too_many_arguments)]
2269fn dspark_attention(
2270    layer: &memra_gguf::model_plan::LayerPlan,
2271    weights: &ReferenceWeights,
2272    x: &[f32],
2273    ring: &[f32],
2274    start_position: usize,
2275    block_size: usize,
2276    hidden: usize,
2277) -> Result<Vec<f32>, ReferenceError> {
2278    use memra_gguf::dsv4_forward::{ActQuantVariant, apply_rope, matmul, rmsnorm};
2279    use memra_gguf::model_plan::{MlaAttentionPlan, RopeFactors, SparseIndexPlan};
2280
2281    let AttentionPlan::Mla(MlaAttentionPlan::CompressedKv {
2282        query_heads,
2283        q_lora_rank,
2284        latent_head_dim,
2285        rope_head_dim,
2286        output_lora_rank,
2287        output_groups,
2288        window,
2289        rope,
2290        compressor: None,
2291        sparse_index: SparseIndexPlan::None,
2292    }) = &layer.attention
2293    else {
2294        return Err(ReferenceError::InvalidPlan {
2295            layer: Some(layer.index),
2296            reason: "DSpark block requires window-only compressed-attention geometry",
2297        });
2298    };
2299    if !matches!(rope.factors, RopeFactors::None) {
2300        return Err(ReferenceError::InvalidPlan {
2301            layer: Some(layer.index),
2302            reason: "DSpark block RoPE must not use scaling factors",
2303        });
2304    }
2305    let heads = *query_heads as usize;
2306    let q_rank = *q_lora_rank as usize;
2307    let head_dim = *latent_head_dim as usize;
2308    let rope_dim = *rope_head_dim as usize;
2309    let output_rank = *output_lora_rank as usize;
2310    let groups = *output_groups as usize;
2311    let window = *window as usize;
2312    if start_position == 0
2313        || head_dim <= rope_dim
2314        || (head_dim - rope_dim) % 64 != 0
2315        || groups == 0
2316        || heads % groups != 0
2317        || ring.len() != window * head_dim
2318    {
2319        return Err(ReferenceError::InvalidPlan {
2320            layer: Some(layer.index),
2321            reason: "DSpark attention has invalid geometry or unprimed ring",
2322        });
2323    }
2324    let positions: Vec<_> = (1..=block_size)
2325        .map(|offset| start_position + offset)
2326        .collect();
2327    let frequencies = memra_gguf::dsv4_forward::precompute_freqs_cis(
2328        rope_dim,
2329        start_position + block_size + 1,
2330        0,
2331        rope.base,
2332        1.0,
2333        32.0,
2334        1.0,
2335    );
2336    let query_low_rank = rmsnorm(
2337        &matmul(
2338            x,
2339            block_size,
2340            hidden,
2341            tensor(
2342                weights,
2343                &layer_id(layer.index, LayerTensor::MlaQueryDown),
2344                &[q_rank, hidden],
2345            )?,
2346            q_rank,
2347        ),
2348        tensor(
2349            weights,
2350            &layer_id(layer.index, LayerTensor::MlaQueryDownNorm),
2351            &[q_rank],
2352        )?,
2353        layer.pre_attention_norm.epsilon,
2354    );
2355    let mut query = matmul(
2356        &query_low_rank,
2357        block_size,
2358        q_rank,
2359        tensor(
2360            weights,
2361            &layer_id(layer.index, LayerTensor::MlaQueryUp),
2362            &[heads * head_dim, q_rank],
2363        )?,
2364        heads * head_dim,
2365    );
2366    for head in query.chunks_exact_mut(head_dim) {
2367        let mean_square = head
2368            .iter()
2369            .map(|value| (*value as f64) * (*value as f64))
2370            .sum::<f64>()
2371            / head_dim as f64;
2372        let scale = 1.0 / (mean_square as f32 + layer.pre_attention_norm.epsilon).sqrt();
2373        for value in head {
2374            *value *= scale;
2375        }
2376    }
2377    apply_rope(
2378        &mut query,
2379        block_size,
2380        heads,
2381        head_dim,
2382        rope_dim,
2383        &frequencies,
2384        &positions,
2385        false,
2386    );
2387    let mut key_value = rmsnorm(
2388        &matmul(
2389            x,
2390            block_size,
2391            hidden,
2392            tensor(
2393                weights,
2394                &layer_id(layer.index, LayerTensor::MlaKvDown),
2395                &[head_dim, hidden],
2396            )?,
2397            head_dim,
2398        ),
2399        tensor(
2400            weights,
2401            &layer_id(layer.index, LayerTensor::MlaKvDownNorm),
2402            &[head_dim],
2403        )?,
2404        layer.pre_attention_norm.epsilon,
2405    );
2406    apply_rope(
2407        &mut key_value,
2408        block_size,
2409        1,
2410        head_dim,
2411        rope_dim,
2412        &frequencies,
2413        &positions,
2414        false,
2415    );
2416    for row in key_value.chunks_exact_mut(head_dim) {
2417        memra_gguf::dsv4_forward::act_quant(
2418            &mut row[..head_dim - rope_dim],
2419            64,
2420            ActQuantVariant::RefFp8Round,
2421        );
2422    }
2423    let indices = memra_gguf::dsv4_dspark::dspark_topk_idxs(window, block_size, start_position);
2424    let sink = tensor(
2425        weights,
2426        &layer_id(layer.index, LayerTensor::AttentionSink),
2427        &[heads],
2428    )?;
2429    let mut attended = vec![0.0; block_size * heads * head_dim];
2430    for token in 0..block_size {
2431        memra_gguf::dsv4_decode::sparse_attn_query(
2432            &query[token * heads * head_dim..(token + 1) * heads * head_dim],
2433            heads,
2434            head_dim,
2435            &indices,
2436            |index| {
2437                if index < window {
2438                    &ring[index * head_dim..(index + 1) * head_dim]
2439                } else {
2440                    let index = index - window;
2441                    &key_value[index * head_dim..(index + 1) * head_dim]
2442                }
2443            },
2444            sink,
2445            (head_dim as f64).powf(-0.5) as f32,
2446            &mut attended[token * heads * head_dim..(token + 1) * heads * head_dim],
2447        );
2448    }
2449    apply_rope(
2450        &mut attended,
2451        block_size,
2452        heads,
2453        head_dim,
2454        rope_dim,
2455        &frequencies,
2456        &positions,
2457        true,
2458    );
2459    let group_width = heads / groups * head_dim;
2460    let output_down = tensor(
2461        weights,
2462        &layer_id(layer.index, LayerTensor::MlaOutputDown),
2463        &[groups * output_rank, group_width],
2464    )?;
2465    let mut grouped = vec![0.0; block_size * groups * output_rank];
2466    for token in 0..block_size {
2467        for group in 0..groups {
2468            let source = &attended[token * heads * head_dim + group * group_width
2469                ..token * heads * head_dim + (group + 1) * group_width];
2470            for rank in 0..output_rank {
2471                let weight = &output_down[(group * output_rank + rank) * group_width
2472                    ..(group * output_rank + rank + 1) * group_width];
2473                grouped[(token * groups + group) * output_rank + rank] =
2474                    memra_gguf::dsv4_forward::dot(source, weight);
2475            }
2476        }
2477    }
2478    Ok(matmul(
2479        &grouped,
2480        block_size,
2481        groups * output_rank,
2482        tensor(
2483            weights,
2484            &layer_id(layer.index, LayerTensor::MlaOutput),
2485            &[hidden, groups * output_rank],
2486        )?,
2487        hidden,
2488    ))
2489}
2490
2491fn hyper_topology(plan: &ModelPlan) -> Result<Option<(usize, f32, u32)>, ReferenceError> {
2492    let topology = plan.layers.iter().find_map(|layer| match layer.residual {
2493        ResidualTopology::HyperConnections {
2494            streams,
2495            epsilon,
2496            sinkhorn_iterations,
2497        } => Some((streams as usize, epsilon, sinkhorn_iterations)),
2498        _ => None,
2499    });
2500    let Some(topology) = topology else {
2501        return Ok(None);
2502    };
2503    if topology.0 == 0 || topology.1 <= 0.0 || topology.2 == 0 {
2504        return Err(ReferenceError::InvalidPlan {
2505            layer: None,
2506            reason: "HyperConnections require streams, epsilon, and Sinkhorn iterations",
2507        });
2508    }
2509    for layer in &plan.layers {
2510        if layer.residual
2511            != (ResidualTopology::HyperConnections {
2512                streams: topology.0 as u32,
2513                epsilon: topology.1,
2514                sinkhorn_iterations: topology.2,
2515            })
2516        {
2517            return Err(ReferenceError::InvalidPlan {
2518                layer: Some(layer.index),
2519                reason: "HyperConnections topology must be consistent across the trunk",
2520            });
2521        }
2522    }
2523    Ok(Some(topology))
2524}
2525
2526fn collapse_hyper_head(
2527    weights: &ReferenceWeights,
2528    x: &[f32],
2529    tokens: usize,
2530    streams: usize,
2531    hidden: usize,
2532    plan: &ModelPlan,
2533    epsilon: f32,
2534) -> Result<Vec<f32>, ReferenceError> {
2535    let set = memra_gguf::dsv4_forward::HcSet {
2536        rows: streams,
2537        fn_w: tensor(
2538            weights,
2539            &TensorId::HyperHeadFunction,
2540            &[streams, streams * hidden],
2541        )?
2542        .to_vec(),
2543        base: tensor(weights, &TensorId::HyperHeadBase, &[streams])?.to_vec(),
2544        scale: tensor(weights, &TensorId::HyperHeadScale, &[1])?.to_vec(),
2545    };
2546    Ok(memra_gguf::dsv4_forward::hc_head(
2547        x,
2548        tokens,
2549        streams,
2550        hidden,
2551        &set,
2552        plan.output_norm.epsilon,
2553        epsilon,
2554    ))
2555}
2556
2557fn apply_logits_transforms(logits: &mut [f32], vocab: usize, transforms: &[LogitsTransform]) {
2558    for transform in transforms {
2559        match transform {
2560            LogitsTransform::Softcap(cap) => {
2561                for value in logits.iter_mut() {
2562                    *value = *cap * (*value / *cap).tanh();
2563                }
2564            }
2565            LogitsTransform::SuppressTokens(ids) => {
2566                for row in logits.chunks_exact_mut(vocab) {
2567                    for &id in ids {
2568                        if let Some(value) = row.get_mut(id as usize) {
2569                            *value = f32::NEG_INFINITY;
2570                        }
2571                    }
2572                }
2573            }
2574        }
2575    }
2576}
2577
2578fn execute_layer(
2579    layer: &memra_gguf::model_plan::LayerPlan,
2580    weights: &ReferenceWeights,
2581    input: &[f32],
2582    token_ids: &[u32],
2583    tokens: usize,
2584    hidden: usize,
2585    vocab: usize,
2586) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
2587    if let ResidualTopology::HyperConnections {
2588        streams,
2589        epsilon,
2590        sinkhorn_iterations,
2591    } = layer.residual
2592    {
2593        return execute_hyper_layer(
2594            layer,
2595            weights,
2596            input,
2597            token_ids,
2598            tokens,
2599            hidden,
2600            vocab,
2601            streams as usize,
2602            epsilon,
2603            sinkhorn_iterations,
2604        );
2605    }
2606    if let ResidualTopology::Gemma {
2607        parallel_moe: Some(parallel),
2608        ..
2609    } = layer.residual
2610    {
2611        return execute_gemma_parallel_moe_layer(layer, parallel, weights, input, tokens, hidden);
2612    }
2613    if let ResidualTopology::Gemma {
2614        parallel_moe: None, ..
2615    } = layer.residual
2616    {
2617        return execute_gemma_dense_layer(layer, weights, input, tokens, hidden);
2618    }
2619    if layer.residual != ResidualTopology::Serial {
2620        return Err(ReferenceError::UnsupportedOperation {
2621            layer: Some(layer.index),
2622            operation: "non-serial residual",
2623        });
2624    }
2625    let pre_attn = rms_norm(
2626        input,
2627        tokens,
2628        hidden,
2629        tensor(
2630            weights,
2631            &layer_id(layer.index, LayerTensor::PreAttentionNorm),
2632            &[hidden],
2633        )?,
2634        layer.pre_attention_norm.epsilon,
2635    );
2636    let (attention, layer_state) = match &layer.attention {
2637        AttentionPlan::Full(attention) => full_attention(
2638            layer.index,
2639            attention,
2640            None,
2641            layer.pre_attention_norm.epsilon,
2642            weights,
2643            &pre_attn,
2644            tokens,
2645            hidden,
2646        )?,
2647        AttentionPlan::SlidingWindow { attention, window } => full_attention(
2648            layer.index,
2649            attention,
2650            Some(*window as usize),
2651            layer.pre_attention_norm.epsilon,
2652            weights,
2653            &pre_attn,
2654            tokens,
2655            hidden,
2656        )?,
2657        AttentionPlan::Mla(mla) => mla_attention(
2658            layer.index,
2659            mla,
2660            layer.pre_attention_norm.epsilon,
2661            weights,
2662            &pre_attn,
2663            tokens,
2664            hidden,
2665        )?,
2666        AttentionPlan::GatedDeltaNet(gdn) => gated_delta_net(
2667            layer.index,
2668            gdn,
2669            layer.pre_attention_norm.epsilon,
2670            weights,
2671            &pre_attn,
2672            tokens,
2673            hidden,
2674        )?,
2675    };
2676    let mut output = input.to_vec();
2677    add_in_place(&mut output, &attention);
2678    let pre_mlp = rms_norm(
2679        &output,
2680        tokens,
2681        hidden,
2682        tensor(
2683            weights,
2684            &layer_id(layer.index, LayerTensor::PreMlpNorm),
2685            &[hidden],
2686        )?,
2687        layer.pre_mlp_norm.epsilon,
2688    );
2689    let mlp = match &layer.mlp {
2690        MlpPlan::Dense(mlp) => dense_mlp(layer.index, mlp, weights, &pre_mlp, tokens, hidden)?,
2691        MlpPlan::Moe(moe) => moe_mlp(
2692            layer.index,
2693            moe,
2694            weights,
2695            &pre_mlp,
2696            token_ids,
2697            tokens,
2698            hidden,
2699            vocab,
2700        )?,
2701    };
2702    add_in_place(&mut output, &mlp);
2703    Ok((output, layer_state))
2704}
2705
2706fn execute_gemma_parallel_moe_layer(
2707    layer: &memra_gguf::model_plan::LayerPlan,
2708    parallel: memra_gguf::model_plan::GemmaParallelMoePlan,
2709    weights: &ReferenceWeights,
2710    input: &[f32],
2711    tokens: usize,
2712    hidden: usize,
2713) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
2714    let ResidualTopology::Gemma {
2715        post_attention_norm,
2716        post_mlp_norm,
2717        layer_scale,
2718        parallel_moe: Some(_),
2719    } = layer.residual
2720    else {
2721        unreachable!()
2722    };
2723    let pre_attention = rms_norm(
2724        input,
2725        tokens,
2726        hidden,
2727        tensor(
2728            weights,
2729            &layer_id(layer.index, LayerTensor::PreAttentionNorm),
2730            &[hidden],
2731        )?,
2732        layer.pre_attention_norm.epsilon,
2733    );
2734    let (attention, state) = match &layer.attention {
2735        AttentionPlan::Full(attention) => full_attention(
2736            layer.index,
2737            attention,
2738            None,
2739            layer.pre_attention_norm.epsilon,
2740            weights,
2741            &pre_attention,
2742            tokens,
2743            hidden,
2744        )?,
2745        AttentionPlan::SlidingWindow { attention, window } => full_attention(
2746            layer.index,
2747            attention,
2748            Some(*window as usize),
2749            layer.pre_attention_norm.epsilon,
2750            weights,
2751            &pre_attention,
2752            tokens,
2753            hidden,
2754        )?,
2755        _ => {
2756            return Err(ReferenceError::UnsupportedOperation {
2757                layer: Some(layer.index),
2758                operation: "gemma parallel MoE non-softmax attention",
2759            });
2760        }
2761    };
2762    let attention = rms_norm(
2763        &attention,
2764        tokens,
2765        hidden,
2766        tensor(
2767            weights,
2768            &layer_id(layer.index, LayerTensor::PostAttentionNorm),
2769            &[hidden],
2770        )?,
2771        post_attention_norm.epsilon,
2772    );
2773    let mut attention_residual = input.to_vec();
2774    add_in_place(&mut attention_residual, &attention);
2775
2776    let MlpPlan::Moe(moe) = &layer.mlp else {
2777        return Err(ReferenceError::InvalidPlan {
2778            layer: Some(layer.index),
2779            reason: "gemma parallel MoE residual requires an MoE plan",
2780        });
2781    };
2782    let shared_plan = moe.shared.as_ref().ok_or(ReferenceError::InvalidPlan {
2783        layer: Some(layer.index),
2784        reason: "gemma parallel MoE requires a shared MLP branch",
2785    })?;
2786    let shared_input = rms_norm(
2787        &attention_residual,
2788        tokens,
2789        hidden,
2790        tensor(
2791            weights,
2792            &layer_id(layer.index, LayerTensor::PreMlpNorm),
2793            &[hidden],
2794        )?,
2795        layer.pre_mlp_norm.epsilon,
2796    );
2797    let shared_intermediate = shared_plan.intermediate_size as usize;
2798    let shared_gate = linear(
2799        &shared_input,
2800        tensor(
2801            weights,
2802            &layer_id(layer.index, LayerTensor::SharedMlpGate),
2803            &[shared_intermediate, hidden],
2804        )?,
2805        tokens,
2806        hidden,
2807        shared_intermediate,
2808    );
2809    let shared_up = linear(
2810        &shared_input,
2811        tensor(
2812            weights,
2813            &layer_id(layer.index, LayerTensor::SharedMlpUp),
2814            &[shared_intermediate, hidden],
2815        )?,
2816        tokens,
2817        hidden,
2818        shared_intermediate,
2819    );
2820    let mut shared_activated = vec![0.0; shared_gate.len()];
2821    for index in 0..shared_activated.len() {
2822        shared_activated[index] = activate_pair(
2823            &moe.activation,
2824            shared_gate[index],
2825            shared_up[index],
2826            layer.index,
2827        )?;
2828    }
2829    let shared = linear(
2830        &shared_activated,
2831        tensor(
2832            weights,
2833            &layer_id(layer.index, LayerTensor::SharedMlpDown),
2834            &[hidden, shared_intermediate],
2835        )?,
2836        tokens,
2837        shared_intermediate,
2838        hidden,
2839    );
2840    let shared = rms_norm(
2841        &shared,
2842        tokens,
2843        hidden,
2844        tensor(
2845            weights,
2846            &layer_id(layer.index, LayerTensor::PostSharedMlpNorm),
2847            &[hidden],
2848        )?,
2849        parallel.shared_post_norm.epsilon,
2850    );
2851
2852    let routed_input = rms_norm(
2853        &attention_residual,
2854        tokens,
2855        hidden,
2856        tensor(
2857            weights,
2858            &layer_id(layer.index, LayerTensor::PreRoutedMlpNorm),
2859            &[hidden],
2860        )?,
2861        parallel.routed_pre_norm.epsilon,
2862    );
2863    let router_scale = tensor(
2864        weights,
2865        &layer_id(layer.index, LayerTensor::MoeRouterScale),
2866        &[hidden],
2867    )?;
2868    let router_weight: Vec<_> = router_scale
2869        .iter()
2870        .map(|value| *value / (hidden as f32).sqrt())
2871        .collect();
2872    let router_input = rms_norm(
2873        &attention_residual,
2874        tokens,
2875        hidden,
2876        &router_weight,
2877        layer.pre_mlp_norm.epsilon,
2878    );
2879    let experts = moe.expert_count as usize;
2880    let selected = moe.experts_per_token as usize;
2881    let intermediate = moe.expert_intermediate_size as usize;
2882    let router_logits = linear(
2883        &router_input,
2884        tensor(
2885            weights,
2886            &layer_id(layer.index, LayerTensor::MoeRouter),
2887            &[experts, hidden],
2888        )?,
2889        tokens,
2890        hidden,
2891        experts,
2892    );
2893    let gate_up = tensor(
2894        weights,
2895        &layer_id(layer.index, LayerTensor::MoeExpertGateUpBank),
2896        &[experts, 2 * intermediate, hidden],
2897    )?;
2898    let down = tensor(
2899        weights,
2900        &layer_id(layer.index, LayerTensor::MoeExpertDownBank),
2901        &[experts, hidden, intermediate],
2902    )?;
2903    let expert_scale = tensor(
2904        weights,
2905        &layer_id(layer.index, LayerTensor::MoeExpertOutputScale),
2906        &[experts],
2907    )?;
2908    let mut routed = vec![0.0; tokens * hidden];
2909    for token in 0..tokens {
2910        let routes = route_experts(
2911            &moe.router,
2912            &router_logits[token * experts..(token + 1) * experts],
2913            None,
2914            selected,
2915            None,
2916            layer.index,
2917        )?;
2918        let row = &routed_input[token * hidden..(token + 1) * hidden];
2919        for (expert, route_weight) in routes {
2920            let expert_offset = expert * 2 * intermediate * hidden;
2921            let mut activated = vec![0.0; intermediate];
2922            for output in 0..intermediate {
2923                let gate = memra_gguf::dsv4_forward::dot(
2924                    row,
2925                    &gate_up
2926                        [expert_offset + output * hidden..expert_offset + (output + 1) * hidden],
2927                );
2928                let up_offset = expert_offset + (intermediate + output) * hidden;
2929                let up =
2930                    memra_gguf::dsv4_forward::dot(row, &gate_up[up_offset..up_offset + hidden]);
2931                activated[output] = activate_pair(&moe.activation, gate, up, layer.index)?;
2932            }
2933            let down_offset = expert * hidden * intermediate;
2934            for output in 0..hidden {
2935                routed[token * hidden + output] += route_weight
2936                    * expert_scale[expert]
2937                    * memra_gguf::dsv4_forward::dot(
2938                        &activated,
2939                        &down[down_offset + output * intermediate
2940                            ..down_offset + (output + 1) * intermediate],
2941                    );
2942            }
2943        }
2944    }
2945    let routed = rms_norm(
2946        &routed,
2947        tokens,
2948        hidden,
2949        tensor(
2950            weights,
2951            &layer_id(layer.index, LayerTensor::PostRoutedMlpNorm),
2952            &[hidden],
2953        )?,
2954        parallel.routed_post_norm.epsilon,
2955    );
2956    let mut combined = shared;
2957    add_in_place(&mut combined, &routed);
2958    let combined = rms_norm(
2959        &combined,
2960        tokens,
2961        hidden,
2962        tensor(
2963            weights,
2964            &layer_id(layer.index, LayerTensor::PostMlpNorm),
2965            &[hidden],
2966        )?,
2967        post_mlp_norm.epsilon,
2968    );
2969    add_in_place(&mut attention_residual, &combined);
2970    let scale = match layer_scale {
2971        GemmaLayerScale::Learned => tensor(
2972            weights,
2973            &layer_id(layer.index, LayerTensor::LayerScale),
2974            &[1],
2975        )?[0],
2976    };
2977    for value in &mut attention_residual {
2978        *value *= scale;
2979    }
2980    Ok((attention_residual, state))
2981}
2982
2983#[allow(clippy::too_many_arguments)]
2984fn execute_hyper_layer(
2985    layer: &memra_gguf::model_plan::LayerPlan,
2986    weights: &ReferenceWeights,
2987    input: &[f32],
2988    token_ids: &[u32],
2989    tokens: usize,
2990    hidden: usize,
2991    vocab: usize,
2992    streams: usize,
2993    epsilon: f32,
2994    sinkhorn_iterations: u32,
2995) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
2996    if input.len() != tokens * streams * hidden {
2997        return Err(ReferenceError::InvalidPlan {
2998            layer: Some(layer.index),
2999            reason: "HyperConnections input does not match tokens x streams x hidden",
3000        });
3001    }
3002    let attention_set = hyper_set(
3003        weights,
3004        layer.index,
3005        streams,
3006        hidden,
3007        LayerTensor::HyperAttentionFunction,
3008        LayerTensor::HyperAttentionBase,
3009        LayerTensor::HyperAttentionScale,
3010    )?;
3011    let (attention_input, post, combination) = memra_gguf::dsv4_forward::hc_pre(
3012        input,
3013        tokens,
3014        streams,
3015        hidden,
3016        &attention_set,
3017        sinkhorn_iterations,
3018        epsilon,
3019    );
3020    let attention_input = rms_norm(
3021        &attention_input,
3022        tokens,
3023        hidden,
3024        tensor(
3025            weights,
3026            &layer_id(layer.index, LayerTensor::PreAttentionNorm),
3027            &[hidden],
3028        )?,
3029        layer.pre_attention_norm.epsilon,
3030    );
3031    let (attention, state) = match &layer.attention {
3032        AttentionPlan::Full(attention) => full_attention(
3033            layer.index,
3034            attention,
3035            None,
3036            layer.pre_attention_norm.epsilon,
3037            weights,
3038            &attention_input,
3039            tokens,
3040            hidden,
3041        )?,
3042        AttentionPlan::SlidingWindow { attention, window } => full_attention(
3043            layer.index,
3044            attention,
3045            Some(*window as usize),
3046            layer.pre_attention_norm.epsilon,
3047            weights,
3048            &attention_input,
3049            tokens,
3050            hidden,
3051        )?,
3052        AttentionPlan::Mla(mla) => mla_attention(
3053            layer.index,
3054            mla,
3055            layer.pre_attention_norm.epsilon,
3056            weights,
3057            &attention_input,
3058            tokens,
3059            hidden,
3060        )?,
3061        AttentionPlan::GatedDeltaNet(gdn) => gated_delta_net(
3062            layer.index,
3063            gdn,
3064            layer.pre_attention_norm.epsilon,
3065            weights,
3066            &attention_input,
3067            tokens,
3068            hidden,
3069        )?,
3070    };
3071    let attention_residual = memra_gguf::dsv4_forward::hc_post(
3072        &attention,
3073        input,
3074        tokens,
3075        streams,
3076        hidden,
3077        &post,
3078        &combination,
3079    );
3080
3081    let mlp_set = hyper_set(
3082        weights,
3083        layer.index,
3084        streams,
3085        hidden,
3086        LayerTensor::HyperMlpFunction,
3087        LayerTensor::HyperMlpBase,
3088        LayerTensor::HyperMlpScale,
3089    )?;
3090    let (mlp_input, post, combination) = memra_gguf::dsv4_forward::hc_pre(
3091        &attention_residual,
3092        tokens,
3093        streams,
3094        hidden,
3095        &mlp_set,
3096        sinkhorn_iterations,
3097        epsilon,
3098    );
3099    let mlp_input = rms_norm(
3100        &mlp_input,
3101        tokens,
3102        hidden,
3103        tensor(
3104            weights,
3105            &layer_id(layer.index, LayerTensor::PreMlpNorm),
3106            &[hidden],
3107        )?,
3108        layer.pre_mlp_norm.epsilon,
3109    );
3110    let mlp = match &layer.mlp {
3111        MlpPlan::Dense(mlp) => dense_mlp(layer.index, mlp, weights, &mlp_input, tokens, hidden)?,
3112        MlpPlan::Moe(moe) => moe_mlp(
3113            layer.index,
3114            moe,
3115            weights,
3116            &mlp_input,
3117            token_ids,
3118            tokens,
3119            hidden,
3120            vocab,
3121        )?,
3122    };
3123    let output = memra_gguf::dsv4_forward::hc_post(
3124        &mlp,
3125        &attention_residual,
3126        tokens,
3127        streams,
3128        hidden,
3129        &post,
3130        &combination,
3131    );
3132    Ok((output, state))
3133}
3134
3135#[allow(clippy::too_many_arguments)]
3136fn hyper_set(
3137    weights: &ReferenceWeights,
3138    layer: u32,
3139    streams: usize,
3140    hidden: usize,
3141    function: LayerTensor,
3142    base: LayerTensor,
3143    scale: LayerTensor,
3144) -> Result<memra_gguf::dsv4_forward::HcSet, ReferenceError> {
3145    let rows = (2 + streams) * streams;
3146    Ok(memra_gguf::dsv4_forward::HcSet {
3147        rows,
3148        fn_w: tensor(
3149            weights,
3150            &layer_id(layer, function),
3151            &[rows, streams * hidden],
3152        )?
3153        .to_vec(),
3154        base: tensor(weights, &layer_id(layer, base), &[rows])?.to_vec(),
3155        scale: tensor(weights, &layer_id(layer, scale), &[3])?.to_vec(),
3156    })
3157}
3158
3159fn execute_gemma_dense_layer(
3160    layer: &memra_gguf::model_plan::LayerPlan,
3161    weights: &ReferenceWeights,
3162    input: &[f32],
3163    tokens: usize,
3164    hidden: usize,
3165) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
3166    let ResidualTopology::Gemma {
3167        post_attention_norm,
3168        post_mlp_norm,
3169        layer_scale,
3170        parallel_moe: None,
3171    } = layer.residual
3172    else {
3173        return Err(ReferenceError::UnsupportedOperation {
3174            layer: Some(layer.index),
3175            operation: "gemma parallel MoE residual",
3176        });
3177    };
3178    let pre_attn = rms_norm(
3179        input,
3180        tokens,
3181        hidden,
3182        tensor(
3183            weights,
3184            &layer_id(layer.index, LayerTensor::PreAttentionNorm),
3185            &[hidden],
3186        )?,
3187        layer.pre_attention_norm.epsilon,
3188    );
3189    let (attention, state) = match &layer.attention {
3190        AttentionPlan::Full(attention) => full_attention(
3191            layer.index,
3192            attention,
3193            None,
3194            layer.pre_attention_norm.epsilon,
3195            weights,
3196            &pre_attn,
3197            tokens,
3198            hidden,
3199        )?,
3200        AttentionPlan::SlidingWindow { attention, window } => full_attention(
3201            layer.index,
3202            attention,
3203            Some(*window as usize),
3204            layer.pre_attention_norm.epsilon,
3205            weights,
3206            &pre_attn,
3207            tokens,
3208            hidden,
3209        )?,
3210        _ => {
3211            return Err(ReferenceError::UnsupportedOperation {
3212                layer: Some(layer.index),
3213                operation: "gemma non-softmax attention",
3214            });
3215        }
3216    };
3217    let post_attention = rms_norm(
3218        &attention,
3219        tokens,
3220        hidden,
3221        tensor(
3222            weights,
3223            &layer_id(layer.index, LayerTensor::PostAttentionNorm),
3224            &[hidden],
3225        )?,
3226        post_attention_norm.epsilon,
3227    );
3228    let mut attention_residual = input.to_vec();
3229    add_in_place(&mut attention_residual, &post_attention);
3230    let pre_mlp = rms_norm(
3231        &attention_residual,
3232        tokens,
3233        hidden,
3234        tensor(
3235            weights,
3236            &layer_id(layer.index, LayerTensor::PreMlpNorm),
3237            &[hidden],
3238        )?,
3239        layer.pre_mlp_norm.epsilon,
3240    );
3241    let MlpPlan::Dense(mlp) = &layer.mlp else {
3242        return Err(ReferenceError::UnsupportedOperation {
3243            layer: Some(layer.index),
3244            operation: "gemma parallel MoE residual",
3245        });
3246    };
3247    let mlp = dense_mlp(layer.index, mlp, weights, &pre_mlp, tokens, hidden)?;
3248    let mlp = rms_norm(
3249        &mlp,
3250        tokens,
3251        hidden,
3252        tensor(
3253            weights,
3254            &layer_id(layer.index, LayerTensor::PostMlpNorm),
3255            &[hidden],
3256        )?,
3257        post_mlp_norm.epsilon,
3258    );
3259    let scale = match layer_scale {
3260        GemmaLayerScale::Learned => tensor(
3261            weights,
3262            &layer_id(layer.index, LayerTensor::LayerScale),
3263            &[1],
3264        )?[0],
3265    };
3266    let mut output = attention_residual;
3267    add_in_place(&mut output, &mlp);
3268    for value in &mut output {
3269        *value *= scale;
3270    }
3271    Ok((output, state))
3272}
3273
3274#[allow(clippy::too_many_arguments)]
3275fn execute_mtp(
3276    plan: &ModelPlan,
3277    weights: &ReferenceWeights,
3278    token_ids: &[u32],
3279    embedding: &[f32],
3280    trunk_hidden: &[f32],
3281    tokens: usize,
3282    hidden: usize,
3283    vocab: usize,
3284    model_output: &[f32],
3285) -> Result<Vec<ReferenceMtpOutput>, ReferenceError> {
3286    if plan.mtp_blocks.is_empty() {
3287        return Ok(Vec::new());
3288    }
3289    if trunk_hidden.len() != tokens * hidden {
3290        return Err(ReferenceError::UnsupportedOperation {
3291            layer: None,
3292            operation: "HyperConnections MTP fusion",
3293        });
3294    }
3295    let mut embedded = vec![0.0; tokens * hidden];
3296    for (position, &token) in token_ids.iter().enumerate() {
3297        let token = token as usize;
3298        embedded[position * hidden..(position + 1) * hidden]
3299            .copy_from_slice(&embedding[token * hidden..(token + 1) * hidden]);
3300    }
3301    let mut source_hidden = trunk_hidden.to_vec();
3302    let mut outputs = Vec::with_capacity(plan.mtp_blocks.len());
3303    for block in &plan.mtp_blocks {
3304        let embedding_norm = rms_norm(
3305            &embedded,
3306            tokens,
3307            hidden,
3308            tensor(
3309                weights,
3310                &TensorId::Mtp {
3311                    depth: block.depth,
3312                    tensor: MtpTensor::EmbeddingNorm,
3313                },
3314                &[hidden],
3315            )?,
3316            block.input.embedding_norm.epsilon,
3317        );
3318        let hidden_norm = rms_norm(
3319            &source_hidden,
3320            tokens,
3321            hidden,
3322            tensor(
3323                weights,
3324                &TensorId::Mtp {
3325                    depth: block.depth,
3326                    tensor: MtpTensor::HiddenNorm,
3327                },
3328                &[hidden],
3329            )?,
3330            block.input.hidden_norm.epsilon,
3331        );
3332        let mut concatenated = vec![0.0; tokens * 2 * hidden];
3333        for token in 0..tokens {
3334            concatenated[token * 2 * hidden..token * 2 * hidden + hidden]
3335                .copy_from_slice(&embedding_norm[token * hidden..(token + 1) * hidden]);
3336            concatenated[token * 2 * hidden + hidden..(token + 1) * 2 * hidden]
3337                .copy_from_slice(&hidden_norm[token * hidden..(token + 1) * hidden]);
3338        }
3339        let fused = linear(
3340            &concatenated,
3341            tensor(
3342                weights,
3343                &TensorId::Mtp {
3344                    depth: block.depth,
3345                    tensor: MtpTensor::FusionProjection,
3346                },
3347                &[hidden, 2 * hidden],
3348            )?,
3349            tokens,
3350            2 * hidden,
3351            hidden,
3352        );
3353        let (hidden_next, state) = execute_layer(
3354            &block.layer,
3355            weights,
3356            &fused,
3357            token_ids,
3358            tokens,
3359            hidden,
3360            vocab,
3361        )?;
3362        let norm_id = TensorId::Mtp {
3363            depth: block.depth,
3364            tensor: MtpTensor::OutputNorm,
3365        };
3366        let norm = match weights.get(&norm_id) {
3367            Some(tensor) => tensor_checked(&norm_id, tensor, &[hidden])?,
3368            None => tensor(weights, &TensorId::OutputNorm, &[hidden])?,
3369        };
3370        let final_hidden = rms_norm(&hidden_next, tokens, hidden, norm, plan.output_norm.epsilon);
3371        let head_id = TensorId::Mtp {
3372            depth: block.depth,
3373            tensor: MtpTensor::OutputProjection,
3374        };
3375        let head = match weights.get(&head_id) {
3376            Some(tensor) => tensor_checked(&head_id, tensor, &[vocab, hidden])?,
3377            None => model_output,
3378        };
3379        let mut logits = linear(&final_hidden, head, tokens, hidden, vocab);
3380        apply_logits_transforms(&mut logits, vocab, &plan.logits);
3381        source_hidden = hidden_next.clone();
3382        outputs.push(ReferenceMtpOutput {
3383            depth: block.depth,
3384            logits,
3385            hidden: hidden_next,
3386            state,
3387        });
3388    }
3389    Ok(outputs)
3390}
3391
3392fn mla_attention(
3393    layer: u32,
3394    plan: &memra_gguf::model_plan::MlaAttentionPlan,
3395    epsilon: f32,
3396    weights: &ReferenceWeights,
3397    x: &[f32],
3398    tokens: usize,
3399    hidden: usize,
3400) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
3401    if let memra_gguf::model_plan::MlaAttentionPlan::CompressedKv { .. } = plan {
3402        return compressed_mla_attention(layer, plan, epsilon, weights, x, tokens, hidden);
3403    }
3404    let memra_gguf::model_plan::MlaAttentionPlan::LatentKv {
3405        query_heads,
3406        q_lora_rank,
3407        kv_lora_rank,
3408        qk_head_dim,
3409        rope_head_dim,
3410        value_head_dim,
3411        rope,
3412        sparse_index,
3413    } = plan.clone()
3414    else {
3415        return Err(ReferenceError::UnsupportedOperation {
3416            layer: Some(layer),
3417            operation: "compressed-KV MLA",
3418        });
3419    };
3420    let sparse_top_k = match sparse_index {
3421        memra_gguf::model_plan::SparseIndexPlan::None => None,
3422        memra_gguf::model_plan::SparseIndexPlan::Own { top_k, .. }
3423        | memra_gguf::model_plan::SparseIndexPlan::SharedFromPrevious { top_k } => {
3424            Some(top_k as usize)
3425        }
3426    };
3427    if sparse_top_k.is_some_and(|top_k| tokens > top_k) {
3428        return Err(ReferenceError::UnsupportedOperation {
3429            layer: Some(layer),
3430            operation: "sparse MLA selection beyond full-selection equivalence",
3431        });
3432    }
3433    let heads = query_heads as usize;
3434    let q_rank = q_lora_rank as usize;
3435    let kv_rank = kv_lora_rank as usize;
3436    let qk_dim = qk_head_dim as usize;
3437    let rope_dim = rope_head_dim as usize;
3438    let nope_dim = qk_dim - rope_dim;
3439    let value_dim = value_head_dim as usize;
3440    let latent_dim = kv_rank + rope_dim;
3441
3442    let q_down = linear(
3443        x,
3444        tensor(
3445            weights,
3446            &layer_id(layer, LayerTensor::MlaQueryDown),
3447            &[q_rank, hidden],
3448        )?,
3449        tokens,
3450        hidden,
3451        q_rank,
3452    );
3453    let q_down = rms_norm(
3454        &q_down,
3455        tokens,
3456        q_rank,
3457        tensor(
3458            weights,
3459            &layer_id(layer, LayerTensor::MlaQueryDownNorm),
3460            &[q_rank],
3461        )?,
3462        epsilon,
3463    );
3464    let query = linear(
3465        &q_down,
3466        tensor(
3467            weights,
3468            &layer_id(layer, LayerTensor::MlaQueryUp),
3469            &[heads * qk_dim, q_rank],
3470        )?,
3471        tokens,
3472        q_rank,
3473        heads * qk_dim,
3474    );
3475    let latent_raw = linear(
3476        x,
3477        tensor(
3478            weights,
3479            &layer_id(layer, LayerTensor::MlaKvDown),
3480            &[latent_dim, hidden],
3481        )?,
3482        tokens,
3483        hidden,
3484        latent_dim,
3485    );
3486    let kv_norm = tensor(
3487        weights,
3488        &layer_id(layer, LayerTensor::MlaKvDownNorm),
3489        &[kv_rank],
3490    )?;
3491    let mut latent = latent_raw;
3492    for token in 0..tokens {
3493        let offset = token * latent_dim;
3494        let normalized = rms_norm(
3495            &latent[offset..offset + kv_rank],
3496            1,
3497            kv_rank,
3498            kv_norm,
3499            epsilon,
3500        );
3501        latent[offset..offset + kv_rank].copy_from_slice(&normalized);
3502    }
3503
3504    let mut query_nope = vec![0.0; tokens * heads * nope_dim];
3505    let mut query_rope = vec![0.0; tokens * heads * rope_dim];
3506    for token in 0..tokens {
3507        for head in 0..heads {
3508            let source = (token * heads + head) * qk_dim;
3509            let nope_target = (token * heads + head) * nope_dim;
3510            let rope_target = (token * heads + head) * rope_dim;
3511            query_nope[nope_target..nope_target + nope_dim]
3512                .copy_from_slice(&query[source..source + nope_dim]);
3513            query_rope[rope_target..rope_target + rope_dim]
3514                .copy_from_slice(&query[source + nope_dim..source + qk_dim]);
3515        }
3516    }
3517    let rope_factors = rope_factor_values(&rope, weights)?;
3518    apply_rope(
3519        &mut query_rope,
3520        tokens,
3521        heads,
3522        rope_dim,
3523        rope.dimensions as usize,
3524        rope.base,
3525        rope_factors.as_deref(),
3526    );
3527    let mut key_rope = vec![0.0; tokens * rope_dim];
3528    for token in 0..tokens {
3529        key_rope[token * rope_dim..(token + 1) * rope_dim]
3530            .copy_from_slice(&latent[token * latent_dim + kv_rank..(token + 1) * latent_dim]);
3531    }
3532    apply_rope(
3533        &mut key_rope,
3534        tokens,
3535        1,
3536        rope_dim,
3537        rope.dimensions as usize,
3538        rope.base,
3539        rope_factors.as_deref(),
3540    );
3541    for token in 0..tokens {
3542        latent[token * latent_dim + kv_rank..(token + 1) * latent_dim]
3543            .copy_from_slice(&key_rope[token * rope_dim..(token + 1) * rope_dim]);
3544    }
3545
3546    let key_weight = tensor(
3547        weights,
3548        &layer_id(layer, LayerTensor::MlaKeyUp),
3549        &[heads, nope_dim, kv_rank],
3550    )?;
3551    let value_weight = tensor(
3552        weights,
3553        &layer_id(layer, LayerTensor::MlaValueUp),
3554        &[heads, value_dim, kv_rank],
3555    )?;
3556    let mut key_nope = vec![0.0; tokens * heads * nope_dim];
3557    let mut value = vec![0.0; tokens * heads * value_dim];
3558    for token in 0..tokens {
3559        let latent_row = &latent[token * latent_dim..token * latent_dim + kv_rank];
3560        for head in 0..heads {
3561            for out in 0..nope_dim {
3562                for rank in 0..kv_rank {
3563                    key_nope[(token * heads + head) * nope_dim + out] +=
3564                        latent_row[rank] * key_weight[(head * nope_dim + out) * kv_rank + rank];
3565                }
3566            }
3567            for out in 0..value_dim {
3568                for rank in 0..kv_rank {
3569                    value[(token * heads + head) * value_dim + out] +=
3570                        latent_row[rank] * value_weight[(head * value_dim + out) * kv_rank + rank];
3571                }
3572            }
3573        }
3574    }
3575    let mut attended = vec![0.0; tokens * heads * value_dim];
3576    let scale = 1.0 / (qk_dim as f32).sqrt();
3577    for token in 0..tokens {
3578        for head in 0..heads {
3579            let mut scores = Vec::with_capacity(token + 1);
3580            for source in 0..=token {
3581                let mut score = 0.0;
3582                for dim in 0..nope_dim {
3583                    score += query_nope[(token * heads + head) * nope_dim + dim]
3584                        * key_nope[(source * heads + head) * nope_dim + dim];
3585                }
3586                for dim in 0..rope_dim {
3587                    score += query_rope[(token * heads + head) * rope_dim + dim]
3588                        * key_rope[source * rope_dim + dim];
3589                }
3590                scores.push(score * scale);
3591            }
3592            softmax_in_place(&mut scores);
3593            for (source, probability) in scores.into_iter().enumerate() {
3594                for dim in 0..value_dim {
3595                    attended[(token * heads + head) * value_dim + dim] +=
3596                        probability * value[(source * heads + head) * value_dim + dim];
3597                }
3598            }
3599        }
3600    }
3601    let output = linear(
3602        &attended,
3603        tensor(
3604            weights,
3605            &layer_id(layer, LayerTensor::MlaOutput),
3606            &[hidden, heads * value_dim],
3607        )?,
3608        tokens,
3609        heads * value_dim,
3610        hidden,
3611    );
3612    Ok((
3613        output,
3614        ReferenceLayerState::LatentKv {
3615            rows: latent,
3616            tokens,
3617            width: latent_dim,
3618        },
3619    ))
3620}
3621
3622fn compressed_mla_attention(
3623    layer: u32,
3624    plan: &memra_gguf::model_plan::MlaAttentionPlan,
3625    epsilon: f32,
3626    weights: &ReferenceWeights,
3627    x: &[f32],
3628    tokens: usize,
3629    hidden: usize,
3630) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
3631    use memra_gguf::dsv4_forward::{
3632        ActQuantVariant, IndexerW, apply_rope as apply_dsv4_rope, matmul, precompute_freqs_cis,
3633        rmsnorm,
3634    };
3635    use memra_gguf::model_plan::{MlaAttentionPlan, RopeFactors, SparseIndexPlan};
3636
3637    let MlaAttentionPlan::CompressedKv {
3638        query_heads,
3639        q_lora_rank,
3640        latent_head_dim,
3641        rope_head_dim,
3642        output_lora_rank,
3643        output_groups,
3644        window,
3645        rope,
3646        compressor,
3647        sparse_index,
3648    } = plan
3649    else {
3650        unreachable!()
3651    };
3652    let heads = *query_heads as usize;
3653    let q_rank = *q_lora_rank as usize;
3654    let head_dim = *latent_head_dim as usize;
3655    let rope_dim = *rope_head_dim as usize;
3656    let output_rank = *output_lora_rank as usize;
3657    let groups = *output_groups as usize;
3658    let window = *window as usize;
3659    if heads == 0
3660        || q_rank == 0
3661        || head_dim == 0
3662        || rope_dim == 0
3663        || rope_dim > head_dim
3664        || (head_dim - rope_dim) % 64 != 0
3665        || groups == 0
3666        || heads % groups != 0
3667        || window == 0
3668    {
3669        return Err(ReferenceError::InvalidPlan {
3670            layer: Some(layer),
3671            reason: "compressed attention has invalid reference geometry",
3672        });
3673    }
3674    let (original_context, factor, beta_fast, beta_slow) = match rope.factors {
3675        RopeFactors::None => (0, 1.0, 32.0, 1.0),
3676        RopeFactors::Yarn {
3677            factor,
3678            original_context,
3679            beta_fast,
3680            beta_slow,
3681        } => (original_context, factor, beta_fast, beta_slow),
3682        _ => {
3683            return Err(ReferenceError::InvalidPlan {
3684                layer: Some(layer),
3685                reason: "compressed attention requires plain or YaRN RoPE",
3686            });
3687        }
3688    };
3689    let frequencies = precompute_freqs_cis(
3690        rope_dim,
3691        tokens.max(1),
3692        original_context,
3693        rope.base,
3694        factor,
3695        beta_fast,
3696        beta_slow,
3697    );
3698    let positions: Vec<usize> = (0..tokens).collect();
3699
3700    let query_low_rank = rmsnorm(
3701        &matmul(
3702            x,
3703            tokens,
3704            hidden,
3705            tensor(
3706                weights,
3707                &layer_id(layer, LayerTensor::MlaQueryDown),
3708                &[q_rank, hidden],
3709            )?,
3710            q_rank,
3711        ),
3712        tensor(
3713            weights,
3714            &layer_id(layer, LayerTensor::MlaQueryDownNorm),
3715            &[q_rank],
3716        )?,
3717        epsilon,
3718    );
3719    let mut query = matmul(
3720        &query_low_rank,
3721        tokens,
3722        q_rank,
3723        tensor(
3724            weights,
3725            &layer_id(layer, LayerTensor::MlaQueryUp),
3726            &[heads * head_dim, q_rank],
3727        )?,
3728        heads * head_dim,
3729    );
3730    for head in query.chunks_exact_mut(head_dim) {
3731        let mean_square = head
3732            .iter()
3733            .map(|value| (*value as f64) * (*value as f64))
3734            .sum::<f64>()
3735            / head_dim as f64;
3736        let scale = 1.0 / (mean_square as f32 + epsilon).sqrt();
3737        for value in head {
3738            *value *= scale;
3739        }
3740    }
3741    apply_dsv4_rope(
3742        &mut query,
3743        tokens,
3744        heads,
3745        head_dim,
3746        rope_dim,
3747        &frequencies,
3748        &positions,
3749        false,
3750    );
3751
3752    let mut key_value = rmsnorm(
3753        &matmul(
3754            x,
3755            tokens,
3756            hidden,
3757            tensor(
3758                weights,
3759                &layer_id(layer, LayerTensor::MlaKvDown),
3760                &[head_dim, hidden],
3761            )?,
3762            head_dim,
3763        ),
3764        tensor(
3765            weights,
3766            &layer_id(layer, LayerTensor::MlaKvDownNorm),
3767            &[head_dim],
3768        )?,
3769        epsilon,
3770    );
3771    apply_dsv4_rope(
3772        &mut key_value,
3773        tokens,
3774        1,
3775        head_dim,
3776        rope_dim,
3777        &frequencies,
3778        &positions,
3779        false,
3780    );
3781    for row in key_value.chunks_exact_mut(head_dim) {
3782        memra_gguf::dsv4_forward::act_quant(
3783            &mut row[..head_dim - rope_dim],
3784            64,
3785            ActQuantVariant::RefFp8Round,
3786        );
3787    }
3788
3789    let (mut indices, mut slots) = memra_gguf::dsv4_forward::window_topk_idxs(window, tokens);
3790    let mut key_value_rows = tokens;
3791    let mut compressed_tokens = 0;
3792    if let Some(compressor_plan) = compressor {
3793        let ratio = compressor_plan.ratio as usize;
3794        let compressor = reference_compressor(
3795            weights,
3796            layer,
3797            hidden,
3798            head_dim,
3799            ratio,
3800            compressor_plan.latent_dim as usize,
3801            false,
3802        )?;
3803        let (compressed_indices, compressed_slots) = match sparse_index {
3804            SparseIndexPlan::None => {
3805                memra_gguf::dsv4_forward::compress_topk_idxs(ratio, tokens, tokens)
3806            }
3807            SparseIndexPlan::Own {
3808                heads: index_heads,
3809                head_dim: index_dim,
3810                top_k,
3811            } => {
3812                let index_heads = *index_heads as usize;
3813                let index_dim = *index_dim as usize;
3814                if index_dim < rope_dim || index_dim % 32 != 0 || !index_dim.is_power_of_two() {
3815                    return Err(ReferenceError::InvalidPlan {
3816                        layer: Some(layer),
3817                        reason: "compressed sparse index has invalid head geometry",
3818                    });
3819                }
3820                let indexer = IndexerW {
3821                    wq_b: tensor(
3822                        weights,
3823                        &layer_id(layer, LayerTensor::SparseQuery),
3824                        &[index_heads * index_dim, q_rank],
3825                    )?
3826                    .to_vec(),
3827                    weights_proj: tensor(
3828                        weights,
3829                        &layer_id(layer, LayerTensor::SparseProjection),
3830                        &[index_heads, hidden],
3831                    )?
3832                    .to_vec(),
3833                    compressor: reference_compressor(
3834                        weights,
3835                        layer,
3836                        hidden,
3837                        index_dim,
3838                        ratio,
3839                        2 * index_dim,
3840                        true,
3841                    )?,
3842                    heads: index_heads,
3843                    hd: index_dim,
3844                    topk: *top_k as usize,
3845                };
3846                let output = indexer.forward(
3847                    x,
3848                    &query_low_rank,
3849                    tokens,
3850                    hidden,
3851                    q_rank,
3852                    tokens,
3853                    &frequencies,
3854                    rope_dim,
3855                    epsilon,
3856                    ActQuantVariant::RefFp8Round,
3857                    false,
3858                );
3859                (output.idxs, output.slots)
3860            }
3861            SparseIndexPlan::SharedFromPrevious { .. } => {
3862                return Err(ReferenceError::UnsupportedOperation {
3863                    layer: Some(layer),
3864                    operation: "shared compressed sparse-index execution",
3865                });
3866            }
3867        };
3868        if compressed_slots > 0 {
3869            let mut merged = vec![-1; tokens * (slots + compressed_slots)];
3870            for token in 0..tokens {
3871                merged[token * (slots + compressed_slots)
3872                    ..token * (slots + compressed_slots) + slots]
3873                    .copy_from_slice(&indices[token * slots..(token + 1) * slots]);
3874                merged[token * (slots + compressed_slots) + slots
3875                    ..(token + 1) * (slots + compressed_slots)]
3876                    .copy_from_slice(
3877                        &compressed_indices
3878                            [token * compressed_slots..(token + 1) * compressed_slots],
3879                    );
3880            }
3881            indices = merged;
3882            slots += compressed_slots;
3883        }
3884        if let Some((compressed, count)) = compressor.forward(
3885            x,
3886            tokens,
3887            hidden,
3888            &frequencies,
3889            rope_dim,
3890            epsilon,
3891            ActQuantVariant::RefFp8Round,
3892        ) {
3893            key_value.extend_from_slice(&compressed);
3894            key_value_rows += count;
3895            compressed_tokens = count;
3896        }
3897    }
3898
3899    let sink = tensor(
3900        weights,
3901        &layer_id(layer, LayerTensor::AttentionSink),
3902        &[heads],
3903    )?;
3904    let attention_scale = (head_dim as f64).powf(-0.5) as f32;
3905    let mut attended = vec![0.0; tokens * heads * head_dim];
3906    for token in 0..tokens {
3907        let selected = &indices[token * slots..(token + 1) * slots];
3908        memra_gguf::dsv4_decode::sparse_attn_query(
3909            &query[token * heads * head_dim..(token + 1) * heads * head_dim],
3910            heads,
3911            head_dim,
3912            selected,
3913            |index| &key_value[index * head_dim..(index + 1) * head_dim],
3914            sink,
3915            attention_scale,
3916            &mut attended[token * heads * head_dim..(token + 1) * heads * head_dim],
3917        );
3918    }
3919    apply_dsv4_rope(
3920        &mut attended,
3921        tokens,
3922        heads,
3923        head_dim,
3924        rope_dim,
3925        &frequencies,
3926        &positions,
3927        true,
3928    );
3929
3930    let group_width = heads / groups * head_dim;
3931    let output_down = tensor(
3932        weights,
3933        &layer_id(layer, LayerTensor::MlaOutputDown),
3934        &[groups * output_rank, group_width],
3935    )?;
3936    let mut grouped = vec![0.0; tokens * groups * output_rank];
3937    for token in 0..tokens {
3938        for group in 0..groups {
3939            let source = &attended[token * heads * head_dim + group * group_width
3940                ..token * heads * head_dim + (group + 1) * group_width];
3941            let group_weight = &output_down
3942                [group * output_rank * group_width..(group + 1) * output_rank * group_width];
3943            for rank in 0..output_rank {
3944                grouped[(token * groups + group) * output_rank + rank] =
3945                    memra_gguf::dsv4_forward::dot(
3946                        source,
3947                        &group_weight[rank * group_width..(rank + 1) * group_width],
3948                    );
3949            }
3950        }
3951    }
3952    let output = matmul(
3953        &grouped,
3954        tokens,
3955        groups * output_rank,
3956        tensor(
3957            weights,
3958            &layer_id(layer, LayerTensor::MlaOutput),
3959            &[hidden, groups * output_rank],
3960        )?,
3961        hidden,
3962    );
3963    Ok((
3964        output,
3965        ReferenceLayerState::CompressedAttention {
3966            rows: key_value,
3967            tokens: key_value_rows,
3968            width: head_dim,
3969            window,
3970            compressed_tokens,
3971        },
3972    ))
3973}
3974
3975#[allow(clippy::too_many_arguments)]
3976fn reference_compressor(
3977    weights: &ReferenceWeights,
3978    layer: u32,
3979    hidden: usize,
3980    output_dim: usize,
3981    ratio: usize,
3982    latent: usize,
3983    sparse: bool,
3984) -> Result<memra_gguf::dsv4_forward::CompressorW, ReferenceError> {
3985    let (key_value, gate, norm, position) = if sparse {
3986        (
3987            LayerTensor::SparseCompressorKeyValue,
3988            LayerTensor::SparseCompressorGate,
3989            LayerTensor::SparseCompressorNorm,
3990            LayerTensor::SparseCompressorPosition,
3991        )
3992    } else {
3993        (
3994            LayerTensor::KvCompressorKeyValue,
3995            LayerTensor::KvCompressorGate,
3996            LayerTensor::KvCompressorNorm,
3997            LayerTensor::KvCompressorPosition,
3998        )
3999    };
4000    Ok(memra_gguf::dsv4_forward::CompressorW {
4001        ratio,
4002        d: output_dim,
4003        latent,
4004        overlap: ratio == 4,
4005        rotate: sparse,
4006        wkv: tensor(weights, &layer_id(layer, key_value), &[latent, hidden])?.to_vec(),
4007        wgate: tensor(weights, &layer_id(layer, gate), &[latent, hidden])?.to_vec(),
4008        norm_w: tensor(weights, &layer_id(layer, norm), &[output_dim])?.to_vec(),
4009        ape: tensor(weights, &layer_id(layer, position), &[ratio, latent])?.to_vec(),
4010    })
4011}
4012
4013fn gated_delta_net(
4014    layer: u32,
4015    plan: &memra_gguf::model_plan::GatedDeltaNetPlan,
4016    epsilon: f32,
4017    weights: &ReferenceWeights,
4018    x: &[f32],
4019    tokens: usize,
4020    hidden: usize,
4021) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
4022    let key_heads = plan.key_heads as usize;
4023    let value_heads = plan.value_heads as usize;
4024    let key_dim = plan.key_head_dim as usize;
4025    let value_dim = plan.value_head_dim as usize;
4026    let kernel = plan.conv_kernel as usize;
4027    if key_heads == 0 || value_heads == 0 || key_dim == 0 || value_dim == 0 || kernel == 0 {
4028        return Err(ReferenceError::InvalidPlan {
4029            layer: Some(layer),
4030            reason: "GDN dimensions must be positive",
4031        });
4032    }
4033    let key_width = key_heads * key_dim;
4034    let value_width = value_heads * value_dim;
4035    let conv_width = 2 * key_width + value_width;
4036    let qkv = linear(
4037        x,
4038        tensor(
4039            weights,
4040            &layer_id(layer, LayerTensor::GdnQkv),
4041            &[conv_width, hidden],
4042        )?,
4043        tokens,
4044        hidden,
4045        conv_width,
4046    );
4047    let gate = linear(
4048        x,
4049        tensor(
4050            weights,
4051            &layer_id(layer, LayerTensor::GdnGate),
4052            &[value_width, hidden],
4053        )?,
4054        tokens,
4055        hidden,
4056        value_width,
4057    );
4058    let beta_raw = linear(
4059        x,
4060        tensor(
4061            weights,
4062            &layer_id(layer, LayerTensor::GdnBeta),
4063            &[value_heads, hidden],
4064        )?,
4065        tokens,
4066        hidden,
4067        value_heads,
4068    );
4069    let alpha = linear(
4070        x,
4071        tensor(
4072            weights,
4073            &layer_id(layer, LayerTensor::GdnAlpha),
4074            &[value_heads, hidden],
4075        )?,
4076        tokens,
4077        hidden,
4078        value_heads,
4079    );
4080    let conv_weight = tensor(
4081        weights,
4082        &layer_id(layer, LayerTensor::GdnConv1d),
4083        &[conv_width, kernel],
4084    )?;
4085    let mut conv = vec![0.0; tokens * conv_width];
4086    let pad = kernel - 1;
4087    for token in 0..tokens {
4088        for channel in 0..conv_width {
4089            let mut sum = 0.0;
4090            for tap in 0..kernel {
4091                let source = token as isize - pad as isize + tap as isize;
4092                if source >= 0 {
4093                    sum += qkv[source as usize * conv_width + channel]
4094                        * conv_weight[channel * kernel + tap];
4095                }
4096            }
4097            conv[token * conv_width + channel] = silu(sum);
4098        }
4099    }
4100
4101    let mut query = vec![0.0; tokens * value_heads * key_dim];
4102    let mut key = vec![0.0; tokens * value_heads * key_dim];
4103    let mut value = vec![0.0; tokens * value_width];
4104    for token in 0..tokens {
4105        for value_head in 0..value_heads {
4106            let key_head = value_head % key_heads;
4107            let q_source = token * conv_width + key_head * key_dim;
4108            let k_source = token * conv_width + key_width + key_head * key_dim;
4109            let v_source = token * conv_width + 2 * key_width + value_head * value_dim;
4110            let q_target = (token * value_heads + value_head) * key_dim;
4111            let v_target = (token * value_heads + value_head) * value_dim;
4112            query[q_target..q_target + key_dim]
4113                .copy_from_slice(&conv[q_source..q_source + key_dim]);
4114            key[q_target..q_target + key_dim].copy_from_slice(&conv[k_source..k_source + key_dim]);
4115            value[v_target..v_target + value_dim]
4116                .copy_from_slice(&conv[v_source..v_source + value_dim]);
4117        }
4118    }
4119    l2_normalize_rows(&mut query, tokens * value_heads, key_dim, epsilon);
4120    l2_normalize_rows(&mut key, tokens * value_heads, key_dim, epsilon);
4121
4122    let a = tensor(weights, &layer_id(layer, LayerTensor::GdnA), &[value_heads])?;
4123    let dt = tensor(
4124        weights,
4125        &layer_id(layer, LayerTensor::GdnDtBias),
4126        &[value_heads],
4127    )?;
4128    let mut matrix = vec![0.0; value_heads * value_dim * key_dim];
4129    let mut mixed = vec![0.0; tokens * value_width];
4130    let scale = 1.0 / (key_dim as f32).sqrt();
4131    for token in 0..tokens {
4132        for head in 0..value_heads {
4133            let beta = sigmoid(beta_raw[token * value_heads + head]);
4134            let decay = (a[head] * softplus(alpha[token * value_heads + head] + dt[head])).exp();
4135            let q_offset = (token * value_heads + head) * key_dim;
4136            let v_offset = (token * value_heads + head) * value_dim;
4137            let state_offset = head * value_dim * key_dim;
4138            let mut next = matrix[state_offset..state_offset + value_dim * key_dim].to_vec();
4139            for value_index in 0..value_dim {
4140                let row = state_offset + value_index * key_dim;
4141                let mut state_key = 0.0;
4142                for key_index in 0..key_dim {
4143                    state_key += matrix[row + key_index] * key[q_offset + key_index];
4144                }
4145                let delta = (value[v_offset + value_index] - decay * state_key) * beta;
4146                let mut attended = 0.0;
4147                for key_index in 0..key_dim {
4148                    let updated =
4149                        decay * matrix[row + key_index] + key[q_offset + key_index] * delta;
4150                    next[value_index * key_dim + key_index] = updated;
4151                    attended += updated * query[q_offset + key_index];
4152                }
4153                mixed[v_offset + value_index] = attended * scale;
4154            }
4155            matrix[state_offset..state_offset + value_dim * key_dim].copy_from_slice(&next);
4156        }
4157    }
4158
4159    let norm = tensor(
4160        weights,
4161        &layer_id(layer, LayerTensor::GdnNorm),
4162        &[value_dim],
4163    )?;
4164    let normalized = rms_norm(&mixed, tokens * value_heads, value_dim, norm, epsilon);
4165    let mut gated = normalized;
4166    for index in 0..gated.len() {
4167        gated[index] *= silu(gate[index]);
4168    }
4169    let output = linear(
4170        &gated,
4171        tensor(
4172            weights,
4173            &layer_id(layer, LayerTensor::GdnOutput),
4174            &[hidden, value_width],
4175        )?,
4176        tokens,
4177        value_width,
4178        hidden,
4179    );
4180    let mut conv_state = vec![0.0; conv_width * pad];
4181    for channel in 0..conv_width {
4182        for index in 0..pad {
4183            let source = tokens as isize - pad as isize + index as isize;
4184            if source >= 0 {
4185                conv_state[channel * pad + index] = qkv[source as usize * conv_width + channel];
4186            }
4187        }
4188    }
4189    Ok((
4190        output,
4191        ReferenceLayerState::Recurrent {
4192            conv: conv_state,
4193            matrix,
4194            value_heads,
4195            key_head_dim: key_dim,
4196            value_head_dim: value_dim,
4197            conv_width,
4198        },
4199    ))
4200}
4201
4202fn full_attention(
4203    layer: u32,
4204    plan: &memra_gguf::model_plan::FullAttentionPlan,
4205    window: Option<usize>,
4206    norm_epsilon: f32,
4207    weights: &ReferenceWeights,
4208    x: &[f32],
4209    tokens: usize,
4210    hidden: usize,
4211) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
4212    let query_heads = plan.query_heads as usize;
4213    let kv_heads = plan.kv_heads as usize;
4214    let key_dim = plan.key_head_dim as usize;
4215    let value_dim = plan.value_head_dim as usize;
4216    if query_heads == 0 || kv_heads == 0 || query_heads % kv_heads != 0 {
4217        return Err(ReferenceError::InvalidPlan {
4218            layer: Some(layer),
4219            reason: "query heads must be a positive multiple of KV heads",
4220        });
4221    }
4222    let fused = plan.output_gate == AttentionGateKind::FusedQ;
4223    let q_width = query_heads * key_dim;
4224    let q_projection_width = q_width * if fused { 2 } else { 1 };
4225    let k_width = kv_heads * key_dim;
4226    let v_width = kv_heads * value_dim;
4227    let q_weight = tensor(
4228        weights,
4229        &layer_id(layer, LayerTensor::Query),
4230        &[q_projection_width, hidden],
4231    )?;
4232    let k_weight = tensor(
4233        weights,
4234        &layer_id(layer, LayerTensor::Key),
4235        &[k_width, hidden],
4236    )?;
4237    let output_weight = tensor(
4238        weights,
4239        &layer_id(layer, LayerTensor::AttentionOutput),
4240        &[hidden, query_heads * value_dim],
4241    )?;
4242    let q_projected = linear(x, q_weight, tokens, hidden, q_projection_width);
4243    let mut query = vec![0.0; tokens * q_width];
4244    let mut fused_gate = None;
4245    if fused {
4246        let mut gate = vec![0.0; tokens * q_width];
4247        for token in 0..tokens {
4248            for head in 0..query_heads {
4249                let projected = token * q_projection_width + head * 2 * key_dim;
4250                let canonical = (token * query_heads + head) * key_dim;
4251                query[canonical..canonical + key_dim]
4252                    .copy_from_slice(&q_projected[projected..projected + key_dim]);
4253                gate[canonical..canonical + key_dim]
4254                    .copy_from_slice(&q_projected[projected + key_dim..projected + 2 * key_dim]);
4255            }
4256        }
4257        fused_gate = Some(gate);
4258    } else {
4259        query.copy_from_slice(&q_projected);
4260    }
4261    let mut key = linear(x, k_weight, tokens, hidden, k_width);
4262    let mut value = match plan.value_projection {
4263        ValueProjection::Separate => linear(
4264            x,
4265            tensor(
4266                weights,
4267                &layer_id(layer, LayerTensor::Value),
4268                &[v_width, hidden],
4269            )?,
4270            tokens,
4271            hidden,
4272            v_width,
4273        ),
4274        ValueProjection::ReuseKey => {
4275            if value_dim != key_dim {
4276                return Err(ReferenceError::InvalidPlan {
4277                    layer: Some(layer),
4278                    reason: "K-as-V requires equal key/value head widths",
4279                });
4280            }
4281            key.clone()
4282        }
4283    };
4284    apply_optional_head_norm(
4285        weights,
4286        layer_id(layer, LayerTensor::QueryNorm),
4287        &mut query,
4288        tokens * query_heads,
4289        key_dim,
4290        plan.qk_norm,
4291        norm_epsilon,
4292    )?;
4293    if plan.value_norm == ValueNorm::WeightlessRms {
4294        let ones = vec![1.0; value_dim];
4295        value = rms_norm(&value, tokens * kv_heads, value_dim, &ones, norm_epsilon);
4296    }
4297    apply_optional_head_norm(
4298        weights,
4299        layer_id(layer, LayerTensor::KeyNorm),
4300        &mut key,
4301        tokens * kv_heads,
4302        key_dim,
4303        plan.qk_norm,
4304        norm_epsilon,
4305    )?;
4306    let rope_factors = rope_factor_values(&plan.rope, weights)?;
4307    apply_rope(
4308        &mut query,
4309        tokens,
4310        query_heads,
4311        key_dim,
4312        plan.rope.dimensions as usize,
4313        plan.rope.base,
4314        rope_factors.as_deref(),
4315    );
4316    apply_rope(
4317        &mut key,
4318        tokens,
4319        kv_heads,
4320        key_dim,
4321        plan.rope.dimensions as usize,
4322        plan.rope.base,
4323        rope_factors.as_deref(),
4324    );
4325
4326    let mut attended = vec![0.0; tokens * query_heads * value_dim];
4327    let scale = match plan.scale {
4328        AttentionScale::InverseSqrtKeyDim => 1.0 / (key_dim as f32).sqrt(),
4329        AttentionScale::Fixed(scale) => scale,
4330    };
4331    for token in 0..tokens {
4332        for head in 0..query_heads {
4333            let kv_head = head * kv_heads / query_heads;
4334            let first_source = window
4335                .map(|window| (token + 1).saturating_sub(window))
4336                .unwrap_or(0);
4337            let mut scores = Vec::with_capacity(token + 1 - first_source);
4338            for source in first_source..=token {
4339                let mut score = 0.0;
4340                for dim in 0..key_dim {
4341                    score += query[(token * query_heads + head) * key_dim + dim]
4342                        * key[(source * kv_heads + kv_head) * key_dim + dim];
4343                }
4344                scores.push(score * scale);
4345            }
4346            softmax_in_place(&mut scores);
4347            for (offset, probability) in scores.into_iter().enumerate() {
4348                let source = first_source + offset;
4349                for dim in 0..value_dim {
4350                    attended[(token * query_heads + head) * value_dim + dim] +=
4351                        probability * value[(source * kv_heads + kv_head) * value_dim + dim];
4352                }
4353            }
4354        }
4355    }
4356    if let Some(gate) = fused_gate {
4357        for token in 0..tokens {
4358            for head in 0..query_heads {
4359                for dim in 0..value_dim {
4360                    if dim >= key_dim {
4361                        return Err(ReferenceError::InvalidPlan {
4362                            layer: Some(layer),
4363                            reason: "fused attention gate requires value_dim <= key_dim",
4364                        });
4365                    }
4366                    attended[(token * query_heads + head) * value_dim + dim] *=
4367                        sigmoid(gate[(token * query_heads + head) * key_dim + dim]);
4368                }
4369            }
4370        }
4371    } else if plan.output_gate == AttentionGateKind::SeparateHead {
4372        let gate_weight = tensor(
4373            weights,
4374            &layer_id(layer, LayerTensor::AttentionGate),
4375            &[query_heads, hidden],
4376        )?;
4377        let gates = linear(x, gate_weight, tokens, hidden, query_heads);
4378        for token in 0..tokens {
4379            for head in 0..query_heads {
4380                let gate = sigmoid(gates[token * query_heads + head]);
4381                for dim in 0..value_dim {
4382                    attended[(token * query_heads + head) * value_dim + dim] *= gate;
4383                }
4384            }
4385        }
4386    }
4387    let state_start = window
4388        .map(|window| tokens.saturating_sub(window))
4389        .unwrap_or(0);
4390    let state_tokens = tokens - state_start;
4391    let state_key = key[state_start * k_width..].to_vec();
4392    let state_value = value[state_start * v_width..].to_vec();
4393    Ok((
4394        linear(
4395            &attended,
4396            output_weight,
4397            tokens,
4398            query_heads * value_dim,
4399            hidden,
4400        ),
4401        ReferenceLayerState::Kv {
4402            key: state_key,
4403            value: state_value,
4404            tokens: state_tokens,
4405            kv_heads,
4406            key_head_dim: key_dim,
4407            value_head_dim: value_dim,
4408            window,
4409        },
4410    ))
4411}
4412
4413fn dense_mlp(
4414    layer: u32,
4415    plan: &memra_gguf::model_plan::DenseMlpPlan,
4416    weights: &ReferenceWeights,
4417    x: &[f32],
4418    tokens: usize,
4419    hidden: usize,
4420) -> Result<Vec<f32>, ReferenceError> {
4421    let intermediate = plan.intermediate_size as usize;
4422    let gate = linear(
4423        x,
4424        tensor(
4425            weights,
4426            &layer_id(layer, LayerTensor::MlpGate),
4427            &[intermediate, hidden],
4428        )?,
4429        tokens,
4430        hidden,
4431        intermediate,
4432    );
4433    let up = linear(
4434        x,
4435        tensor(
4436            weights,
4437            &layer_id(layer, LayerTensor::MlpUp),
4438            &[intermediate, hidden],
4439        )?,
4440        tokens,
4441        hidden,
4442        intermediate,
4443    );
4444    let mut activated = vec![0.0; gate.len()];
4445    for index in 0..gate.len() {
4446        activated[index] = activate_pair(&plan.activation, gate[index], up[index], layer)?;
4447    }
4448    Ok(linear(
4449        &activated,
4450        tensor(
4451            weights,
4452            &layer_id(layer, LayerTensor::MlpDown),
4453            &[hidden, intermediate],
4454        )?,
4455        tokens,
4456        intermediate,
4457        hidden,
4458    ))
4459}
4460
4461fn moe_mlp(
4462    layer: u32,
4463    plan: &memra_gguf::model_plan::MoeMlpPlan,
4464    weights: &ReferenceWeights,
4465    x: &[f32],
4466    token_ids: &[u32],
4467    tokens: usize,
4468    hidden: usize,
4469    vocab: usize,
4470) -> Result<Vec<f32>, ReferenceError> {
4471    let experts = plan.expert_count as usize;
4472    let selected = plan.experts_per_token as usize;
4473    let intermediate = plan.expert_intermediate_size as usize;
4474    if selected == 0 || selected > experts {
4475        return Err(ReferenceError::InvalidPlan {
4476            layer: Some(layer),
4477            reason: "MoE top-k must be in 1..=expert_count",
4478        });
4479    }
4480    let router = tensor(
4481        weights,
4482        &layer_id(layer, LayerTensor::MoeRouter),
4483        &[experts, hidden],
4484    )?;
4485    let logits = linear(x, router, tokens, hidden, experts);
4486    let bias = if router_has_selection_bias(&plan.router) {
4487        Some(tensor(
4488            weights,
4489            &layer_id(layer, LayerTensor::MoeRouterBias),
4490            &[experts],
4491        )?)
4492    } else {
4493        None
4494    };
4495    let token_to_expert = if matches!(
4496        plan.router,
4497        memra_gguf::model_plan::RouterPlan::TokenIdHash { .. }
4498    ) {
4499        Some(tensor(
4500            weights,
4501            &layer_id(layer, LayerTensor::MoeTokenToExpert),
4502            &[vocab, selected],
4503        )?)
4504    } else {
4505        None
4506    };
4507    let gate_bank = tensor(
4508        weights,
4509        &layer_id(layer, LayerTensor::MoeExpertGateBank),
4510        &[experts, intermediate, hidden],
4511    )?;
4512    let up_bank = tensor(
4513        weights,
4514        &layer_id(layer, LayerTensor::MoeExpertUpBank),
4515        &[experts, intermediate, hidden],
4516    )?;
4517    let down_bank = tensor(
4518        weights,
4519        &layer_id(layer, LayerTensor::MoeExpertDownBank),
4520        &[experts, hidden, intermediate],
4521    )?;
4522    let mut output = vec![0.0; tokens * hidden];
4523    for token in 0..tokens {
4524        let forced_routes = token_to_expert
4525            .map(|table| {
4526                let token_id = token_ids[token] as usize;
4527                &table[token_id * selected..(token_id + 1) * selected]
4528            })
4529            .map(|row| {
4530                row.iter()
4531                    .map(|&value| {
4532                        if !value.is_finite()
4533                            || value < 0.0
4534                            || value.fract() != 0.0
4535                            || value as usize >= experts
4536                        {
4537                            return Err(ReferenceError::InvalidPlan {
4538                                layer: Some(layer),
4539                                reason: "token-id expert table contains an invalid expert id",
4540                            });
4541                        }
4542                        Ok(value as usize)
4543                    })
4544                    .collect::<Result<Vec<_>, _>>()
4545            })
4546            .transpose()?;
4547        let routes = route_experts(
4548            &plan.router,
4549            &logits[token * experts..(token + 1) * experts],
4550            bias,
4551            selected,
4552            forced_routes.as_deref(),
4553            layer,
4554        )?;
4555        let input = &x[token * hidden..(token + 1) * hidden];
4556        for (expert, route_weight) in routes {
4557            let gate_offset = expert * intermediate * hidden;
4558            let down_offset = expert * hidden * intermediate;
4559            let mut activated = vec![0.0; intermediate];
4560            for row in 0..intermediate {
4561                let mut gate = 0.0;
4562                let mut up = 0.0;
4563                for column in 0..hidden {
4564                    gate += input[column] * gate_bank[gate_offset + row * hidden + column];
4565                    up += input[column] * up_bank[gate_offset + row * hidden + column];
4566                }
4567                activated[row] = activate_pair(&plan.activation, gate, up, layer)?;
4568            }
4569            for row in 0..hidden {
4570                let mut value = 0.0;
4571                for column in 0..intermediate {
4572                    value +=
4573                        activated[column] * down_bank[down_offset + row * intermediate + column];
4574                }
4575                output[token * hidden + row] += route_weight * value;
4576            }
4577        }
4578    }
4579
4580    if let Some(shared) = plan.shared.as_ref() {
4581        let intermediate = shared.intermediate_size as usize;
4582        let gate = linear(
4583            x,
4584            tensor(
4585                weights,
4586                &layer_id(layer, LayerTensor::SharedMlpGate),
4587                &[intermediate, hidden],
4588            )?,
4589            tokens,
4590            hidden,
4591            intermediate,
4592        );
4593        let up = linear(
4594            x,
4595            tensor(
4596                weights,
4597                &layer_id(layer, LayerTensor::SharedMlpUp),
4598                &[intermediate, hidden],
4599            )?,
4600            tokens,
4601            hidden,
4602            intermediate,
4603        );
4604        let mut activated = vec![0.0; gate.len()];
4605        for index in 0..gate.len() {
4606            activated[index] = activate_pair(&plan.activation, gate[index], up[index], layer)?;
4607        }
4608        let mut shared_output = linear(
4609            &activated,
4610            tensor(
4611                weights,
4612                &layer_id(layer, LayerTensor::SharedMlpDown),
4613                &[hidden, intermediate],
4614            )?,
4615            tokens,
4616            intermediate,
4617            hidden,
4618        );
4619        if shared.gated {
4620            let gate_weight = tensor(
4621                weights,
4622                &layer_id(layer, LayerTensor::SharedMlpInputGate),
4623                &[hidden],
4624            )?;
4625            for token in 0..tokens {
4626                let mut gate = 0.0;
4627                for column in 0..hidden {
4628                    gate += x[token * hidden + column] * gate_weight[column];
4629                }
4630                let gate = sigmoid(gate);
4631                for column in 0..hidden {
4632                    shared_output[token * hidden + column] *= gate;
4633                }
4634            }
4635        }
4636        add_in_place(&mut output, &shared_output);
4637    }
4638    Ok(output)
4639}
4640
4641fn route_experts(
4642    router: &memra_gguf::model_plan::RouterPlan,
4643    logits: &[f32],
4644    bias: Option<&[f32]>,
4645    selected: usize,
4646    forced_indices: Option<&[usize]>,
4647    layer: u32,
4648) -> Result<Vec<(usize, f32)>, ReferenceError> {
4649    use memra_gguf::model_plan::{RouterPlan, RouterScorePlan};
4650
4651    let mut weights = match router {
4652        RouterPlan::Softmax => {
4653            let mut probabilities = logits.to_vec();
4654            softmax_in_place(&mut probabilities);
4655            probabilities
4656        }
4657        RouterPlan::Sigmoid { .. } => logits.iter().map(|&value| sigmoid(value)).collect(),
4658        RouterPlan::SqrtSoftplus { .. } => {
4659            logits.iter().map(|&value| softplus(value).sqrt()).collect()
4660        }
4661        RouterPlan::TokenIdHash { score, .. } => match score {
4662            RouterScorePlan::Softmax => {
4663                let mut probabilities = logits.to_vec();
4664                softmax_in_place(&mut probabilities);
4665                probabilities
4666            }
4667            RouterScorePlan::Sigmoid => logits.iter().map(|&value| sigmoid(value)).collect(),
4668            RouterScorePlan::SqrtSoftplus => {
4669                logits.iter().map(|&value| softplus(value).sqrt()).collect()
4670            }
4671        },
4672    };
4673    let selection_scores: Vec<f32> = weights
4674        .iter()
4675        .enumerate()
4676        .map(|(index, &weight)| weight + bias.map_or(0.0, |bias| bias[index]))
4677        .collect();
4678    let indices = if let RouterPlan::TokenIdHash { .. } = router {
4679        let Some(forced) = forced_indices else {
4680            return Err(ReferenceError::InvalidPlan {
4681                layer: Some(layer),
4682                reason: "token-id hash router requires a token-to-expert row",
4683            });
4684        };
4685        if forced.len() != selected {
4686            return Err(ReferenceError::InvalidPlan {
4687                layer: Some(layer),
4688                reason: "token-id expert row width does not match MoE top-k",
4689            });
4690        }
4691        let mut seen = std::collections::BTreeSet::new();
4692        for &index in forced {
4693            if index >= logits.len() || !seen.insert(index) {
4694                return Err(ReferenceError::InvalidPlan {
4695                    layer: Some(layer),
4696                    reason: "token-id expert row contains an out-of-range or duplicate expert",
4697                });
4698            }
4699        }
4700        forced.to_vec()
4701    } else {
4702        if forced_indices.is_some() {
4703            return Err(ReferenceError::InvalidPlan {
4704                layer: Some(layer),
4705                reason: "score-selected router received forced expert indices",
4706            });
4707        }
4708        let mut indices: Vec<usize> = (0..logits.len()).collect();
4709        indices.sort_by(|&left, &right| {
4710            selection_scores[right]
4711                .total_cmp(&selection_scores[left])
4712                .then(left.cmp(&right))
4713        });
4714        indices.truncate(selected);
4715        indices
4716    };
4717    let (normalize, scaling) = match router {
4718        RouterPlan::Softmax => (true, 1.0),
4719        RouterPlan::Sigmoid {
4720            normalize_selected,
4721            scaling_factor,
4722            ..
4723        }
4724        | RouterPlan::SqrtSoftplus {
4725            normalize_selected,
4726            scaling_factor,
4727            ..
4728        } => (*normalize_selected, *scaling_factor),
4729        RouterPlan::TokenIdHash {
4730            normalize_selected,
4731            scaling_factor,
4732            ..
4733        } => (*normalize_selected, *scaling_factor),
4734    };
4735    if normalize {
4736        let denominator = indices
4737            .iter()
4738            .map(|&index| weights[index])
4739            .sum::<f32>()
4740            .max(if matches!(router, RouterPlan::Softmax) {
4741                6.103_515_6e-5
4742            } else {
4743                1e-20
4744            });
4745        for weight in &mut weights {
4746            *weight = *weight / denominator * scaling;
4747        }
4748    } else {
4749        for weight in &mut weights {
4750            *weight *= scaling;
4751        }
4752    }
4753    Ok(indices
4754        .into_iter()
4755        .map(|index| (index, weights[index]))
4756        .collect())
4757}
4758
4759fn router_has_selection_bias(router: &memra_gguf::model_plan::RouterPlan) -> bool {
4760    matches!(
4761        router,
4762        memra_gguf::model_plan::RouterPlan::Sigmoid {
4763            selection_bias: true,
4764            ..
4765        } | memra_gguf::model_plan::RouterPlan::SqrtSoftplus {
4766            selection_bias: true,
4767            ..
4768        }
4769    )
4770}
4771
4772fn activate_pair(
4773    activation: &ActivationPlan,
4774    gate: f32,
4775    up: f32,
4776    layer: u32,
4777) -> Result<f32, ReferenceError> {
4778    Ok(match activation {
4779        ActivationPlan::Silu => silu(gate) * up,
4780        ActivationPlan::GeluTanh => gelu_tanh(gate) * up,
4781        ActivationPlan::SwiGluOai { alpha, limit } => {
4782            (gate * sigmoid(*alpha * gate)).min(*limit) * up.clamp(-*limit, *limit)
4783        }
4784        ActivationPlan::SwiGluClamped { limit } => {
4785            silu(gate).min(*limit) * up.clamp(-*limit, *limit)
4786        }
4787        ActivationPlan::Named(_) => {
4788            return Err(ReferenceError::UnsupportedOperation {
4789                layer: Some(layer),
4790                operation: "named MLP activation",
4791            });
4792        }
4793    })
4794}
4795
4796fn tensor<'a>(
4797    weights: &'a ReferenceWeights,
4798    id: &TensorId,
4799    expected: &[usize],
4800) -> Result<&'a [f32], ReferenceError> {
4801    let tensor = weights
4802        .get(id)
4803        .ok_or_else(|| ReferenceError::MissingTensor(id.clone()))?;
4804    tensor_checked(id, tensor, expected)
4805}
4806
4807fn tensor_checked<'a>(
4808    id: &TensorId,
4809    tensor: &'a ReferenceTensor,
4810    expected: &[usize],
4811) -> Result<&'a [f32], ReferenceError> {
4812    if tensor.shape != expected {
4813        return Err(ReferenceError::TensorShape {
4814            id: Some(id.clone()),
4815            expected: expected.to_vec(),
4816            actual_elements: tensor.data.len(),
4817        });
4818    }
4819    Ok(&tensor.data)
4820}
4821
4822fn layer_id(layer: u32, tensor: LayerTensor) -> TensorId {
4823    TensorId::Layer {
4824        index: layer,
4825        tensor,
4826    }
4827}
4828
4829fn linear(x: &[f32], weight: &[f32], rows: usize, input: usize, output: usize) -> Vec<f32> {
4830    let mut result = vec![0.0; rows * output];
4831    for row in 0..rows {
4832        for out in 0..output {
4833            let mut sum = 0.0;
4834            for inner in 0..input {
4835                sum += x[row * input + inner] * weight[out * input + inner];
4836            }
4837            result[row * output + out] = sum;
4838        }
4839    }
4840    result
4841}
4842
4843fn rms_norm(x: &[f32], rows: usize, width: usize, weight: &[f32], epsilon: f32) -> Vec<f32> {
4844    let mut result = vec![0.0; x.len()];
4845    for row in 0..rows {
4846        let input = &x[row * width..(row + 1) * width];
4847        let mean_square = input.iter().map(|value| value * value).sum::<f32>() / width as f32;
4848        let inverse = 1.0 / (mean_square + epsilon).sqrt();
4849        for index in 0..width {
4850            result[row * width + index] = input[index] * inverse * weight[index];
4851        }
4852    }
4853    result
4854}
4855
4856fn l2_normalize_rows(values: &mut [f32], rows: usize, width: usize, epsilon: f32) {
4857    for row in 0..rows {
4858        let offset = row * width;
4859        let sum = values[offset..offset + width]
4860            .iter()
4861            .map(|value| value * value)
4862            .sum::<f32>();
4863        let inverse = 1.0 / (sum + epsilon).sqrt();
4864        for value in &mut values[offset..offset + width] {
4865            *value *= inverse;
4866        }
4867    }
4868}
4869
4870fn apply_optional_head_norm(
4871    weights: &ReferenceWeights,
4872    id: TensorId,
4873    values: &mut [f32],
4874    rows: usize,
4875    width: usize,
4876    presence: memra_gguf::model_plan::TensorPresence,
4877    epsilon: f32,
4878) -> Result<(), ReferenceError> {
4879    let Some(weight) = weights.get(&id) else {
4880        return if presence == memra_gguf::model_plan::TensorPresence::Required {
4881            Err(ReferenceError::MissingTensor(id))
4882        } else {
4883            Ok(())
4884        };
4885    };
4886    let normalized = rms_norm(
4887        values,
4888        rows,
4889        width,
4890        tensor_checked(&id, weight, &[width])?,
4891        epsilon,
4892    );
4893    values.copy_from_slice(&normalized);
4894    Ok(())
4895}
4896
4897fn rope_factor_values(
4898    plan: &memra_gguf::model_plan::RopePlan,
4899    weights: &ReferenceWeights,
4900) -> Result<Option<Vec<f32>>, ReferenceError> {
4901    use memra_gguf::model_plan::RopeFactors;
4902
4903    let width = plan.dimensions as usize / 2;
4904    Ok(match plan.factors {
4905        RopeFactors::None => None,
4906        RopeFactors::PartialRotary { factor } => {
4907            let keep = (width as f32 * factor.clamp(0.0, 1.0)).round() as usize;
4908            Some(
4909                (0..width)
4910                    .map(|index| if index < keep { 1.0 } else { 1.0e30 })
4911                    .collect(),
4912            )
4913        }
4914        RopeFactors::Checkpoint => {
4915            let tensor = weights
4916                .get(&TensorId::RopeFactors)
4917                .ok_or_else(|| ReferenceError::MissingTensor(TensorId::RopeFactors))?;
4918            if tensor.shape.len() != 1 || tensor.data.len() < width {
4919                return Err(ReferenceError::TensorShape {
4920                    id: Some(TensorId::RopeFactors),
4921                    expected: vec![width],
4922                    actual_elements: tensor.data.len(),
4923                });
4924            }
4925            Some(tensor.data[..width].to_vec())
4926        }
4927        RopeFactors::Yarn { .. } => {
4928            return Err(ReferenceError::UnsupportedOperation {
4929                layer: None,
4930                operation: "YaRN on non-compressed attention",
4931            });
4932        }
4933    })
4934}
4935
4936fn apply_rope(
4937    values: &mut [f32],
4938    tokens: usize,
4939    heads: usize,
4940    head_dim: usize,
4941    dimensions: usize,
4942    base: f32,
4943    factors: Option<&[f32]>,
4944) {
4945    let dimensions = dimensions.min(head_dim) / 2 * 2;
4946    let half = dimensions / 2;
4947    for token in 0..tokens {
4948        for head in 0..heads {
4949            let offset = (token * heads + head) * head_dim;
4950            for index in 0..half {
4951                let factor = factors.map_or(1.0, |factors| factors[index]);
4952                let frequency = base.powf(-2.0 * index as f32 / dimensions as f32) / factor;
4953                let angle = token as f32 * frequency;
4954                let (sin, cos) = angle.sin_cos();
4955                let first = values[offset + index];
4956                let second = values[offset + index + half];
4957                values[offset + index] = first * cos - second * sin;
4958                values[offset + index + half] = first * sin + second * cos;
4959            }
4960        }
4961    }
4962}
4963
4964fn softmax_in_place(values: &mut [f32]) {
4965    let max = values.iter().copied().fold(f32::NEG_INFINITY, f32::max);
4966    let mut sum = 0.0;
4967    for value in values.iter_mut() {
4968        *value = (*value - max).exp();
4969        sum += *value;
4970    }
4971    for value in values {
4972        *value /= sum;
4973    }
4974}
4975
4976fn add_in_place(target: &mut [f32], addend: &[f32]) {
4977    for (target, addend) in target.iter_mut().zip(addend) {
4978        *target += addend;
4979    }
4980}
4981
4982fn sigmoid(value: f32) -> f32 {
4983    1.0 / (1.0 + (-value).exp())
4984}
4985
4986fn silu(value: f32) -> f32 {
4987    value * sigmoid(value)
4988}
4989
4990fn softplus(value: f32) -> f32 {
4991    if value > 20.0 {
4992        value
4993    } else {
4994        value.exp().ln_1p()
4995    }
4996}
4997
4998fn gelu_tanh(value: f32) -> f32 {
4999    0.5 * value * (1.0 + (0.797_884_6 * (value + 0.044_715 * value * value * value)).tanh())
5000}
5001
5002#[cfg(test)]
5003mod tests {
5004    use super::*;
5005    use memra_gguf::config::{HfConfig, ModelConfig};
5006
5007    fn weight(shape: &[usize], data: &[f32]) -> ReferenceTensor {
5008        ReferenceTensor::new(shape.to_vec(), data.to_vec()).unwrap()
5009    }
5010
5011    #[test]
5012    fn one_token_dense_plan_matches_hand_derived_logits_and_emits_kv_state() {
5013        let config = ModelConfig::from_hf(&HfConfig::parse(
5014            r#"{"model_type":"qwen3","num_hidden_layers":1,"hidden_size":2,
5015            "num_attention_heads":1,"num_key_value_heads":1,"head_dim":2,
5016            "intermediate_size":2,"vocab_size":3,"max_position_embeddings":8,
5017            "rms_norm_eps":0.000001}"#,
5018        ));
5019        let plan = ModelPlan::compile(&config).unwrap();
5020        let identity = [1.0, 0.0, 0.0, 1.0];
5021        let zero = [0.0; 4];
5022        let mut weights = ReferenceWeights::new();
5023        weights.insert(
5024            TensorId::TokenEmbedding,
5025            weight(&[3, 2], &[1.0, 0.0, 0.0, 1.0, -1.0, 0.0]),
5026        );
5027        weights.insert(TensorId::OutputNorm, weight(&[2], &[1.0, 1.0]));
5028        for tensor in [LayerTensor::PreAttentionNorm, LayerTensor::PreMlpNorm] {
5029            weights.insert(layer_id(0, tensor), weight(&[2], &[1.0, 1.0]));
5030        }
5031        for tensor in [
5032            LayerTensor::Query,
5033            LayerTensor::Key,
5034            LayerTensor::Value,
5035            LayerTensor::AttentionOutput,
5036        ] {
5037            weights.insert(layer_id(0, tensor), weight(&[2, 2], &identity));
5038        }
5039        for tensor in [
5040            LayerTensor::MlpGate,
5041            LayerTensor::MlpUp,
5042            LayerTensor::MlpDown,
5043        ] {
5044            weights.insert(layer_id(0, tensor), weight(&[2, 2], &zero));
5045        }
5046
5047        let output = execute(&plan, &weights, &[0]).unwrap();
5048        let root_two = 2.0f32.sqrt();
5049        assert_eq!((output.tokens, output.vocab), (1, 3));
5050        assert!((output.logits[0] - root_two).abs() < 2e-5);
5051        assert!(output.logits[1].abs() < 2e-5);
5052        assert!((output.logits[2] + root_two).abs() < 2e-5);
5053        let ReferenceLayerState::Kv {
5054            tokens, key, value, ..
5055        } = &output.state.layers[0]
5056        else {
5057            panic!("expected KV state");
5058        };
5059        assert_eq!(*tokens, 1);
5060        assert_eq!(key.len(), 2);
5061        assert_eq!(value.len(), 2);
5062    }
5063
5064    #[test]
5065    fn hyperconnections_execute_stream_state_and_head_collapse() {
5066        let config = ModelConfig::from_hf(&HfConfig::parse(
5067            r#"{"model_type":"qwen3","num_hidden_layers":1,"hidden_size":2,
5068            "num_attention_heads":1,"num_key_value_heads":1,"head_dim":2,
5069            "intermediate_size":2,"vocab_size":3,"max_position_embeddings":8}"#,
5070        ));
5071        let mut plan = ModelPlan::compile(&config).unwrap();
5072        plan.layers[0].residual = ResidualTopology::HyperConnections {
5073            streams: 2,
5074            epsilon: 1e-6,
5075            sinkhorn_iterations: 2,
5076        };
5077        let fixture = deterministic_fixture(&plan).unwrap();
5078        assert_eq!(
5079            fixture.weights[&TensorId::HyperHeadFunction].shape,
5080            vec![2, 4]
5081        );
5082        assert_eq!(
5083            fixture.weights[&layer_id(0, LayerTensor::HyperAttentionFunction)].shape,
5084            vec![8, 4]
5085        );
5086        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
5087        assert!(output.logits.iter().all(|value| value.is_finite()));
5088        assert!(matches!(
5089            output.state.layers[0],
5090            ReferenceLayerState::Kv { .. }
5091        ));
5092    }
5093
5094    #[test]
5095    fn generated_tiny_fixture_is_deterministic_and_executable() {
5096        let config = ModelConfig::from_hf(&HfConfig::parse(
5097            r#"{"model_type":"qwen3","num_hidden_layers":2,"hidden_size":8,
5098            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
5099            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32}"#,
5100        ));
5101        let plan = ModelPlan::compile(&config).unwrap();
5102        let first = deterministic_fixture(&plan).unwrap();
5103        let second = deterministic_fixture(&plan).unwrap();
5104        assert_eq!(first, second);
5105        let output = execute(&plan, &first.weights, &first.token_ids).unwrap();
5106        assert_eq!(output.logits.len(), first.token_ids.len() * 32);
5107        assert!(output.logits.iter().all(|value| value.is_finite()));
5108    }
5109
5110    #[test]
5111    fn qwen35_fixture_executes_mixed_gdn_and_full_attention_state() {
5112        let config = ModelConfig::from_hf(&HfConfig::parse(
5113            r#"{"model_type":"qwen3_5","num_hidden_layers":4,"hidden_size":8,
5114            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
5115            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32,
5116            "rms_norm_eps":0.000001,"full_attention_interval":2,
5117            "linear_conv_kernel_dim":3,"linear_key_head_dim":4,
5118            "linear_value_head_dim":4,"linear_num_key_heads":1,
5119            "linear_num_value_heads":2}"#,
5120        ));
5121        let plan = ModelPlan::compile(&config).unwrap();
5122        let fixture = deterministic_fixture(&plan).unwrap();
5123        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
5124        assert_eq!(output.state.layers.len(), 4);
5125        assert!(matches!(
5126            output.state.layers[0],
5127            ReferenceLayerState::Recurrent { .. }
5128        ));
5129        assert!(matches!(
5130            output.state.layers[1],
5131            ReferenceLayerState::Kv { .. }
5132        ));
5133        assert!(matches!(
5134            output.state.layers[2],
5135            ReferenceLayerState::Recurrent { .. }
5136        ));
5137        assert!(matches!(
5138            output.state.layers[3],
5139            ReferenceLayerState::Kv { .. }
5140        ));
5141        assert!(output.logits.iter().all(|value| value.is_finite()));
5142        assert_eq!(
5143            output.logits[..8]
5144                .iter()
5145                .map(|value| value.to_bits())
5146                .collect::<Vec<_>>(),
5147            vec![
5148                3_182_242_076,
5149                1_053_299_392,
5150                3_199_800_546,
5151                3_198_737_445,
5152                3_180_184_136,
5153                3_187_768_631,
5154                1_057_556_100,
5155                1_035_812_924,
5156            ]
5157        );
5158    }
5159
5160    #[test]
5161    fn router_laws_pin_stable_ties_and_selection_only_bias() {
5162        use memra_gguf::model_plan::{RouterPlan, RouterScorePlan};
5163
5164        assert_eq!(
5165            route_experts(&RouterPlan::Softmax, &[0.0, 0.0, 0.0], None, 2, None, 0,).unwrap(),
5166            vec![(0, 0.5), (1, 0.5)]
5167        );
5168        assert_eq!(
5169            route_experts(
5170                &RouterPlan::Sigmoid {
5171                    normalize_selected: true,
5172                    scaling_factor: 2.0,
5173                    selection_bias: true,
5174                },
5175                &[0.0, 0.0],
5176                Some(&[-1.0, 1.0]),
5177                1,
5178                None,
5179                0,
5180            )
5181            .unwrap(),
5182            vec![(1, 2.0)]
5183        );
5184        assert_eq!(
5185            route_experts(
5186                &RouterPlan::TokenIdHash {
5187                    score: RouterScorePlan::SqrtSoftplus,
5188                    normalize_selected: true,
5189                    scaling_factor: 1.5,
5190                },
5191                &[0.0, 0.0, 0.0],
5192                None,
5193                2,
5194                Some(&[2, 0]),
5195                0,
5196            )
5197            .unwrap(),
5198            vec![(2, 0.75), (0, 0.75)]
5199        );
5200        assert!(matches!(
5201            route_experts(
5202                &RouterPlan::TokenIdHash {
5203                    score: RouterScorePlan::SqrtSoftplus,
5204                    normalize_selected: true,
5205                    scaling_factor: 1.5,
5206                },
5207                &[0.0, 0.0, 0.0],
5208                None,
5209                2,
5210                Some(&[1, 1]),
5211                0,
5212            ),
5213            Err(ReferenceError::InvalidPlan {
5214                reason: "token-id expert row contains an out-of-range or duplicate expert",
5215                ..
5216            })
5217        ));
5218    }
5219
5220    #[test]
5221    fn token_hash_moe_fixture_executes_from_semantic_token_table() {
5222        use memra_gguf::model_plan::{RouterPlan, RouterScorePlan};
5223
5224        let config = ModelConfig::from_hf(&HfConfig::parse(
5225            r#"{"model_type":"qwen3_moe","num_hidden_layers":1,"hidden_size":8,
5226            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
5227            "intermediate_size":16,"vocab_size":16,"max_position_embeddings":32,
5228            "num_experts":4,"num_experts_per_tok":2,"moe_intermediate_size":8}"#,
5229        ));
5230        let mut plan = ModelPlan::compile(&config).unwrap();
5231        let MlpPlan::Moe(moe) = &mut plan.layers[0].mlp else {
5232            unreachable!()
5233        };
5234        moe.router = RouterPlan::TokenIdHash {
5235            score: RouterScorePlan::SqrtSoftplus,
5236            normalize_selected: true,
5237            scaling_factor: 1.5,
5238        };
5239        let fixture = deterministic_fixture(&plan).unwrap();
5240        let table_id = layer_id(0, LayerTensor::MoeTokenToExpert);
5241        assert_eq!(fixture.weights[&table_id].shape, vec![16, 2]);
5242        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
5243        assert!(output.logits.iter().all(|value| value.is_finite()));
5244
5245        let mut alternate = fixture.weights.clone();
5246        alternate.get_mut(&table_id).unwrap().data.fill(3.0);
5247        for row in alternate
5248            .get_mut(&table_id)
5249            .unwrap()
5250            .data
5251            .chunks_exact_mut(2)
5252        {
5253            row[1] = 2.0;
5254        }
5255        let alternate = execute(&plan, &alternate, &fixture.token_ids).unwrap();
5256        assert_ne!(output.logits, alternate.logits);
5257    }
5258
5259    #[test]
5260    fn qwen3_moe_fixture_executes_routed_and_shared_branches() {
5261        let config = ModelConfig::from_hf(&HfConfig::parse(
5262            r#"{"model_type":"qwen3_moe","num_hidden_layers":2,"hidden_size":8,
5263            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
5264            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32,
5265            "num_experts":4,"num_experts_per_tok":2,"moe_intermediate_size":8,
5266            "shared_expert_intermediate_size":8}"#,
5267        ));
5268        let plan = ModelPlan::compile(&config).unwrap();
5269        let fixture = deterministic_fixture(&plan).unwrap();
5270        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
5271        assert!(output.logits.iter().all(|value| value.is_finite()));
5272        assert_eq!(
5273            output.logits[..8]
5274                .iter()
5275                .map(|value| value.to_bits())
5276                .collect::<Vec<_>>(),
5277            vec![
5278                3_205_834_204,
5279                1_034_800_117,
5280                1_053_917_366,
5281                3_190_866_844,
5282                984_171_488,
5283                3_182_514_784,
5284                3_154_736_064,
5285                3_175_624_690,
5286            ]
5287        );
5288    }
5289
5290    #[test]
5291    fn sliding_window_limits_attention_and_trims_reference_state() {
5292        let config = ModelConfig::from_hf(&HfConfig::parse(
5293            r#"{"model_type":"qwen3","num_hidden_layers":1,"hidden_size":8,
5294            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
5295            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32}"#,
5296        ));
5297        let mut plan = ModelPlan::compile(&config).unwrap();
5298        let AttentionPlan::Full(attention) = plan.layers[0].attention.clone() else {
5299            unreachable!()
5300        };
5301        plan.layers[0].attention = AttentionPlan::SlidingWindow {
5302            attention,
5303            window: 2,
5304        };
5305        let fixture = deterministic_fixture(&plan).unwrap();
5306        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
5307        let ReferenceLayerState::Kv { tokens, window, .. } = output.state.layers[0] else {
5308            panic!("expected sliding KV state");
5309        };
5310        assert_eq!(tokens, 2);
5311        assert_eq!(window, Some(2));
5312    }
5313
5314    #[test]
5315    fn mla_fixture_emits_latent_state_and_sparse_overflow_refuses() {
5316        use memra_gguf::model_plan::{
5317            MlaAttentionPlan, RopeFactors, RopePlan, SparseIndexPlan, StatePlan,
5318        };
5319
5320        let config = ModelConfig::from_hf(&HfConfig::parse(
5321            r#"{"model_type":"qwen3","num_hidden_layers":1,"hidden_size":8,
5322            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
5323            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32}"#,
5324        ));
5325        let mut plan = ModelPlan::compile(&config).unwrap();
5326        let mla = MlaAttentionPlan::LatentKv {
5327            query_heads: 2,
5328            q_lora_rank: 4,
5329            kv_lora_rank: 4,
5330            qk_head_dim: 4,
5331            rope_head_dim: 2,
5332            value_head_dim: 4,
5333            rope: RopePlan {
5334                dimensions: 2,
5335                base: 10_000.0,
5336                factors: RopeFactors::None,
5337            },
5338            sparse_index: SparseIndexPlan::None,
5339        };
5340        plan.layers[0].attention = AttentionPlan::Mla(mla.clone());
5341        plan.layers[0].state = StatePlan::LatentKvCache { width: 6 };
5342        let fixture = deterministic_fixture(&plan).unwrap();
5343        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
5344        let ReferenceLayerState::LatentKv { tokens, width, .. } = output.state.layers[0] else {
5345            panic!("expected latent KV state");
5346        };
5347        assert_eq!((tokens, width), (3, 6));
5348        assert_eq!(
5349            output.logits[..4]
5350                .iter()
5351                .map(|value| value.to_bits())
5352                .collect::<Vec<_>>(),
5353            vec![1_035_177_220, 1_055_447_641, 3_201_478_680, 3_199_508_856]
5354        );
5355
5356        let MlaAttentionPlan::LatentKv {
5357            query_heads,
5358            q_lora_rank,
5359            kv_lora_rank,
5360            qk_head_dim,
5361            rope_head_dim,
5362            value_head_dim,
5363            rope,
5364            ..
5365        } = mla
5366        else {
5367            unreachable!()
5368        };
5369        plan.layers[0].attention = AttentionPlan::Mla(MlaAttentionPlan::LatentKv {
5370            query_heads,
5371            q_lora_rank,
5372            kv_lora_rank,
5373            qk_head_dim,
5374            rope_head_dim,
5375            value_head_dim,
5376            rope,
5377            sparse_index: SparseIndexPlan::Own {
5378                heads: 1,
5379                head_dim: 2,
5380                top_k: 2,
5381            },
5382        });
5383        let error = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap_err();
5384        assert!(matches!(
5385            error,
5386            ReferenceError::UnsupportedOperation {
5387                operation: "sparse MLA selection beyond full-selection equivalence",
5388                ..
5389            }
5390        ));
5391    }
5392
5393    #[test]
5394    fn compressed_mla_executes_window_compressor_indexer_and_grouped_output() {
5395        use memra_gguf::model_plan::{
5396            KvCompressorPlan, MlaAttentionPlan, RopeFactors, RopePlan, SparseIndexPlan, StatePlan,
5397        };
5398
5399        let config = ModelConfig::from_hf(&HfConfig::parse(
5400            r#"{"model_type":"qwen3","num_hidden_layers":1,"hidden_size":128,
5401            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":64,
5402            "intermediate_size":256,"vocab_size":32,"max_position_embeddings":64,
5403            "rms_norm_eps":0.000001}"#,
5404        ));
5405        let mut plan = ModelPlan::compile(&config).unwrap();
5406        plan.layers[0].attention = AttentionPlan::Mla(MlaAttentionPlan::CompressedKv {
5407            query_heads: 2,
5408            q_lora_rank: 64,
5409            latent_head_dim: 128,
5410            rope_head_dim: 64,
5411            output_lora_rank: 64,
5412            output_groups: 1,
5413            window: 4,
5414            rope: RopePlan {
5415                dimensions: 64,
5416                base: 160_000.0,
5417                factors: RopeFactors::Yarn {
5418                    factor: 2.0,
5419                    original_context: 32,
5420                    beta_fast: 32.0,
5421                    beta_slow: 1.0,
5422                },
5423            },
5424            compressor: Some(KvCompressorPlan {
5425                ratio: 4,
5426                latent_dim: 256,
5427            }),
5428            sparse_index: SparseIndexPlan::Own {
5429                heads: 2,
5430                head_dim: 128,
5431                top_k: 2,
5432            },
5433        });
5434        plan.layers[0].state = StatePlan::CompressedAttention {
5435            window: 4,
5436            head_dim: 128,
5437            compressor_ratio: Some(4),
5438            sparse_top_k: Some(2),
5439        };
5440        let fixture = deterministic_fixture(&plan).unwrap();
5441        let output = execute(&plan, &fixture.weights, &[1, 2, 3, 4]).unwrap();
5442        let ReferenceLayerState::CompressedAttention {
5443            tokens,
5444            width,
5445            window,
5446            compressed_tokens,
5447            ..
5448        } = output.state.layers[0]
5449        else {
5450            panic!("expected compressed attention state")
5451        };
5452        assert_eq!((tokens, width, window, compressed_tokens), (5, 128, 4, 1));
5453        assert!(output.logits.iter().all(|value| value.is_finite()));
5454    }
5455
5456    #[test]
5457    fn dsv4_shaped_trunk_executes_one_canonical_plan() {
5458        let config = ModelConfig::from_hf(&HfConfig::parse(
5459            r#"{"model_type":"deepseek_v4","num_hidden_layers":2,"hidden_size":128,
5460            "num_attention_heads":1,"num_key_value_heads":1,"head_dim":128,
5461            "intermediate_size":256,"vocab_size":128,"max_position_embeddings":1024,
5462            "rms_norm_eps":0.000001,"rope_theta":10000,"n_routed_experts":4,
5463            "n_shared_experts":1,"num_experts_per_tok":2,"moe_intermediate_size":128,
5464            "norm_topk_prob":true,"num_hash_layers":1,"num_nextn_predict_layers":1,
5465            "scoring_func":"sqrtsoftplus","topk_method":"noaux_tc",
5466            "routed_scaling_factor":1.5,"hc_eps":0.000001,"hc_mult":2,
5467            "hc_sinkhorn_iters":4,"q_lora_rank":128,"qk_rope_head_dim":64,
5468            "o_lora_rank":128,"o_groups":1,"index_n_heads":1,"index_head_dim":128,
5469            "index_topk":16,"compress_ratios":[0,4,0],"compress_rope_theta":160000,
5470            "sliding_window":128,"swiglu_limit":10.0,
5471            "rope_scaling":{"factor":4,"beta_fast":32,"beta_slow":1,
5472            "original_max_position_embeddings":1024}}"#,
5473        ));
5474        let mut plan = ModelPlan::compile(&config).unwrap();
5475        assert_eq!(plan.layers.len(), 2);
5476        plan.mtp_blocks.clear();
5477        let fixture = deterministic_fixture(&plan).unwrap();
5478        let output = execute(&plan, &fixture.weights, &[1, 2, 3, 4]).unwrap();
5479        assert_eq!(output.state.layers.len(), 2);
5480        assert!(
5481            output
5482                .state
5483                .layers
5484                .iter()
5485                .all(|state| matches!(state, ReferenceLayerState::CompressedAttention { .. }))
5486        );
5487        assert!(
5488            fixture
5489                .weights
5490                .contains_key(&layer_id(0, LayerTensor::MoeTokenToExpert))
5491        );
5492        assert!(
5493            fixture
5494                .weights
5495                .contains_key(&layer_id(1, LayerTensor::MoeRouterBias))
5496        );
5497        assert!(output.logits.iter().all(|value| value.is_finite()));
5498    }
5499
5500    #[test]
5501    fn dspark_executes_trunk_tap_ring_blocks_markov_and_confidence() {
5502        use memra_gguf::model_plan::{DrafterPlan, DsparkPlan};
5503
5504        let config = ModelConfig::from_hf(&HfConfig::parse(
5505            r#"{"model_type":"deepseek_v4","num_hidden_layers":2,"hidden_size":128,
5506            "num_attention_heads":1,"num_key_value_heads":1,"head_dim":128,
5507            "intermediate_size":256,"vocab_size":128,"max_position_embeddings":1024,
5508            "rms_norm_eps":0.000001,"rope_theta":10000,"n_routed_experts":4,
5509            "n_shared_experts":1,"num_experts_per_tok":2,"moe_intermediate_size":128,
5510            "norm_topk_prob":true,"num_hash_layers":1,"num_nextn_predict_layers":1,
5511            "scoring_func":"sqrtsoftplus","topk_method":"noaux_tc",
5512            "routed_scaling_factor":1.5,"hc_eps":0.000001,"hc_mult":2,
5513            "hc_sinkhorn_iters":4,"q_lora_rank":128,"qk_rope_head_dim":64,
5514            "o_lora_rank":128,"o_groups":1,"index_n_heads":1,"index_head_dim":128,
5515            "index_topk":16,"compress_ratios":[0,4,0],"compress_rope_theta":160000,
5516            "sliding_window":128,"swiglu_limit":10.0,
5517            "rope_scaling":{"factor":4,"beta_fast":32,"beta_slow":1,
5518            "original_max_position_embeddings":1024}}"#,
5519        ));
5520        let mut plan = ModelPlan::compile(&config).unwrap();
5521        let block = plan.mtp_blocks.remove(0).layer;
5522        plan.drafter = Some(DrafterPlan::Dspark(DsparkPlan {
5523            block_size: 3,
5524            noise_token_id: 31,
5525            target_layer_ids: vec![1],
5526            markov_rank: 8,
5527            blocks: vec![block],
5528        }));
5529        let fixture = deterministic_fixture(&plan).unwrap();
5530        let output = execute(&plan, &fixture.weights, &[1, 2, 3, 4]).unwrap();
5531        let draft = output.draft.expect("DSpark output");
5532        assert_eq!(draft.input_token, 4);
5533        assert_eq!(draft.output_ids.len(), 4);
5534        assert_eq!(draft.confidence.len(), 3);
5535        assert_eq!(draft.logits.len(), 3 * 128);
5536        assert!(draft.logits.iter().all(|value| value.is_finite()));
5537        assert!(draft.confidence.iter().all(|value| value.is_finite()));
5538    }
5539
5540    #[test]
5541    fn gemma4_vision_executes_patch_rope_pool_standardize_and_projection() {
5542        let config = ModelConfig::from_hf(&HfConfig::parse(
5543            r#"{"model_type":"gemma4","image_token_id":31,"vision_soft_tokens_per_image":1,
5544            "text_config":{"model_type":"gemma4_text",
5545            "num_hidden_layers":2,"hidden_size":8,"num_attention_heads":2,
5546            "num_key_value_heads":1,"num_global_key_value_heads":1,"head_dim":4,
5547            "global_head_dim":4,"intermediate_size":16,"vocab_size":32,
5548            "max_position_embeddings":64,"rms_norm_eps":0.000001,"sliding_window":8,
5549            "layer_types":["sliding_attention","full_attention"],
5550            "rope_parameters":{"full_attention":{"rope_theta":10000,
5551            "partial_rotary_factor":0.5},"sliding_attention":{"rope_theta":10000}}},
5552            "vision_config":{"hidden_size":8,"intermediate_size":16,
5553            "num_hidden_layers":2,"num_attention_heads":2,"num_key_value_heads":1,
5554            "head_dim":4,"max_position_embeddings":64,"patch_size":2,
5555            "position_embedding_size":16,"pooling_kernel_size":2,
5556            "rms_norm_eps":0.000001,"standardize":true,"use_clipped_linears":false,
5557            "hidden_activation":"gelu_pytorch_tanh","rope_parameters":{"rope_theta":100}}}"#,
5558        ));
5559        let plan = ModelPlan::compile(&config).unwrap();
5560        let fixture = deterministic_fixture(&plan).unwrap();
5561        let input = fixture.vision.as_ref().expect("vision fixture");
5562        let first = execute_vision(&plan, &fixture.weights, input).unwrap();
5563        let second = execute_vision(&plan, &fixture.weights, input).unwrap();
5564        assert_eq!(first, second);
5565        assert_eq!((first.patch_count, first.output_tokens), (4, 1));
5566        assert_eq!((first.hidden_size, first.projection_size), (8, 8));
5567        assert_eq!(first.encoder_hidden.len(), 4 * 8);
5568        assert_eq!(first.pooled_hidden.len(), 8);
5569        assert_eq!(first.projected_hidden.len(), 8);
5570        assert!(first.projected_hidden.iter().all(|value| value.is_finite()));
5571        let multimodal = execute_multimodal(&plan, &fixture.weights, &[1, 31, 2], input).unwrap();
5572        let text_only = execute(&plan, &fixture.weights, &[1, 31, 2]).unwrap();
5573        assert_eq!(multimodal.vision, first);
5574        assert_ne!(multimodal.language.logits, text_only.logits);
5575        assert!(
5576            plan.operations()
5577                .contains(&memra_gguf::model_plan::OperationKind::VisionTokenInjection)
5578        );
5579    }
5580
5581    #[test]
5582    fn gemma4_parallel_moe_executes_shared_routed_and_scaled_residual_branches() {
5583        let config = ModelConfig::from_hf(&HfConfig::parse(
5584            r#"{"model_type":"gemma4","text_config":{"model_type":"gemma4_text",
5585            "num_hidden_layers":2,"hidden_size":8,"num_attention_heads":2,
5586            "num_key_value_heads":1,"num_global_key_value_heads":1,"head_dim":4,
5587            "global_head_dim":4,"intermediate_size":16,"moe_intermediate_size":8,
5588            "num_experts":4,"top_k_experts":2,"vocab_size":32,
5589            "max_position_embeddings":64,"rms_norm_eps":0.000001,"sliding_window":8,
5590            "layer_types":["sliding_attention","full_attention"],
5591            "rope_parameters":{"full_attention":{"rope_theta":10000,
5592            "partial_rotary_factor":0.5},"sliding_attention":{"rope_theta":10000}}}}"#,
5593        ));
5594        let plan = ModelPlan::compile(&config).unwrap();
5595        let MlpPlan::Moe(moe) = &plan.layers[0].mlp else {
5596            panic!("expected Gemma MoE")
5597        };
5598        assert_eq!(moe.experts_per_token, 2);
5599        assert_eq!(moe.shared.as_ref().unwrap().intermediate_size, 16);
5600        assert!(matches!(
5601            plan.layers[0].residual,
5602            ResidualTopology::Gemma {
5603                parallel_moe: Some(_),
5604                ..
5605            }
5606        ));
5607        let fixture = deterministic_fixture(&plan).unwrap();
5608        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
5609        assert!(output.logits.iter().all(|value| value.is_finite()));
5610        assert!(
5611            plan.operations()
5612                .contains(&memra_gguf::model_plan::OperationKind::GemmaParallelMoeResidual)
5613        );
5614    }
5615
5616    #[test]
5617    fn embedded_mtp_executes_typed_fusion_block_and_fallback_head() {
5618        let config = ModelConfig::from_hf(&HfConfig::parse(
5619            r#"{"model_type":"qwen3_5","num_hidden_layers":2,
5620            "num_nextn_predict_layers":1,"hidden_size":8,
5621            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
5622            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32,
5623            "rms_norm_eps":0.000001,"full_attention_interval":2,
5624            "linear_conv_kernel_dim":3,"linear_key_head_dim":4,
5625            "linear_value_head_dim":4,"linear_num_key_heads":1,
5626            "linear_num_value_heads":2}"#,
5627        ));
5628        let plan = ModelPlan::compile(&config).unwrap();
5629        assert_eq!(plan.mtp_blocks.len(), 1);
5630        let fixture = deterministic_fixture(&plan).unwrap();
5631        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
5632        assert_eq!(output.mtp.len(), 1);
5633        assert_eq!(output.mtp[0].depth, 0);
5634        assert_eq!(output.mtp[0].hidden.len(), fixture.token_ids.len() * 8);
5635        assert_eq!(output.mtp[0].logits.len(), fixture.token_ids.len() * 32);
5636        assert!(output.mtp[0].logits.iter().all(|value| value.is_finite()));
5637        assert_eq!(
5638            output.mtp[0].logits[..4]
5639                .iter()
5640                .map(|value| value.to_bits())
5641                .collect::<Vec<_>>(),
5642            vec![1_042_962_358, 1_044_718_512, 3_171_782_004, 3_189_261_409]
5643        );
5644    }
5645
5646    #[test]
5647    fn multi_depth_mtp_threads_hidden_through_every_typed_block() {
5648        let config = ModelConfig::from_hf(&HfConfig::parse(
5649            r#"{"model_type":"qwen3_5","num_hidden_layers":2,
5650            "num_nextn_predict_layers":2,"hidden_size":8,
5651            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
5652            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32,
5653            "rms_norm_eps":0.000001,"full_attention_interval":2,
5654            "linear_conv_kernel_dim":3,"linear_key_head_dim":4,
5655            "linear_value_head_dim":4,"linear_num_key_heads":1,
5656            "linear_num_value_heads":2}"#,
5657        ));
5658        let plan = ModelPlan::compile(&config).unwrap();
5659        assert_eq!(plan.mtp_blocks.len(), 2);
5660        let fixture = deterministic_fixture(&plan).unwrap();
5661        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
5662        assert_eq!(
5663            output
5664                .mtp
5665                .iter()
5666                .map(|block| block.depth)
5667                .collect::<Vec<_>>(),
5668            vec![0, 1]
5669        );
5670        assert!(
5671            output
5672                .mtp
5673                .iter()
5674                .flat_map(|block| &block.logits)
5675                .all(|value| value.is_finite())
5676        );
5677        assert_ne!(output.mtp[0].hidden, output.mtp[1].hidden);
5678    }
5679
5680    #[test]
5681    fn rope_uses_neox_split_half_pairs() {
5682        use memra_gguf::model_plan::{RopeFactors, RopePlan};
5683
5684        let mut values = vec![1.0, 2.0, 3.0, 4.0];
5685        apply_rope(&mut values, 1, 1, 4, 4, 10_000.0, None);
5686        // Position zero is deliberately unchanged.
5687        assert_eq!(values, vec![1.0, 2.0, 3.0, 4.0]);
5688
5689        let mut values = vec![0.0; 8];
5690        values[4..].copy_from_slice(&[1.0, 2.0, 3.0, 4.0]);
5691        apply_rope(&mut values, 2, 1, 4, 4, 10_000.0, None);
5692        let (sin0, cos0) = 1.0f32.sin_cos();
5693        let (sin1, cos1) = 0.01f32.sin_cos();
5694        let row = &values[4..];
5695        assert!((row[0] - (cos0 - 3.0 * sin0)).abs() < 1e-6);
5696        assert!((row[2] - (sin0 + 3.0 * cos0)).abs() < 1e-6);
5697        assert!((row[1] - (2.0 * cos1 - 4.0 * sin1)).abs() < 1e-6);
5698        assert!((row[3] - (2.0 * sin1 + 4.0 * cos1)).abs() < 1e-6);
5699        assert_eq!(
5700            rope_factor_values(
5701                &RopePlan {
5702                    dimensions: 4,
5703                    base: 10_000.0,
5704                    factors: RopeFactors::PartialRotary { factor: 0.5 },
5705                },
5706                &ReferenceWeights::new(),
5707            )
5708            .unwrap()
5709            .unwrap(),
5710            vec![1.0, 1.0e30]
5711        );
5712    }
5713
5714    #[test]
5715    fn dense_gemma_executes_scaled_parallel_residual_and_k_as_v() {
5716        let config = ModelConfig::from_hf(&HfConfig::parse(
5717            r#"{"model_type":"gemma4","num_hidden_layers":2,"hidden_size":8,
5718            "num_attention_heads":2,"num_key_value_heads":1,
5719            "num_global_key_value_heads":1,"head_dim":4,"global_head_dim":4,
5720            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32,
5721            "rms_norm_eps":0.000001,"sliding_window":2,
5722            "final_logit_softcapping":30,
5723            "layer_types":["sliding_attention","full_attention"],
5724            "rope_parameters":{"full_attention":{"rope_theta":1000000,
5725            "partial_rotary_factor":0.5},"sliding_attention":{"rope_theta":10000}}}"#,
5726        ));
5727        let plan = ModelPlan::compile(&config).unwrap();
5728        assert_eq!(plan.embedding_scale, 8.0f32.sqrt());
5729        let fixture = deterministic_fixture(&plan).unwrap();
5730        assert!(
5731            !fixture
5732                .weights
5733                .contains_key(&layer_id(1, LayerTensor::Value))
5734        );
5735        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
5736        assert!(output.logits.iter().all(|value| value.is_finite()));
5737        let ReferenceLayerState::Kv { window, .. } = output.state.layers[0] else {
5738            panic!("expected SWA state");
5739        };
5740        assert_eq!(window, Some(2));
5741        let ReferenceLayerState::Kv { window, .. } = output.state.layers[1] else {
5742            panic!("expected global state");
5743        };
5744        assert_eq!(window, None);
5745        assert_eq!(
5746            output.logits[..4]
5747                .iter()
5748                .map(|value| value.to_bits())
5749                .collect::<Vec<_>>(),
5750            vec![3_198_203_366, 1_057_194_687, 3_185_247_713, 3_204_119_266]
5751        );
5752    }
5753}