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