1use 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 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 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 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 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}