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