1use crate::bundle::{load_bundle, Bundle, LoadedWeight};
2use crate::chat::{apply_chat_template, strip_assistant_visible, ChatTurn};
3use crate::family::{
4 effective_rope_theta, graph_hook, infer_family_path, require_runnable, ArchClass, Family,
5};
6use crate::multimodal::asr_transcribe_pcm16le;
7use crate::profile::{
8 elapsed_ms, load_profile_begin, load_profile_set_cuda_upload, load_profile_set_materialize,
9 load_profile_set_mmap, load_profile_take, EngineProfile, GenerateProfile,
10};
11use crate::tensor_names::{
12 action_head_names, altup_projection_names, altup_unembed_names, attn_k_names,
13 attn_k_norm_names, attn_norm_names, attn_o_names, attn_post_norm_names, attn_q_names,
14 attn_q_norm_names, attn_v_names, attn_v_norm_names, conv_in_proj_names, conv_kernel_names,
15 conv_out_proj_names, emb_names, embed_per_layer_names, ffn_down_names, ffn_gate_names,
16 ffn_norm_names, ffn_post_norm_names, ffn_up_names, layer_altup_correct_scale_names,
17 layer_altup_correction_coef_names, layer_altup_prediction_coef_names, layer_altup_router_names,
18 layer_altup_router_norm_names, layer_laurel_left_names, layer_laurel_norm_names,
19 layer_laurel_right_names, layer_ple_gate_names, layer_ple_post_norm_names,
20 layer_ple_proj_names, layer_scalar_names, linear_a_log_names, linear_conv1d_names,
21 linear_dt_bias_names, linear_in_proj_a_names, linear_in_proj_b_names, linear_in_proj_ba_names,
22 linear_in_proj_qkv_names, linear_in_proj_qkvz_names, linear_in_proj_z_names,
23 linear_out_norm_names, linear_out_proj_names, moe_expert_down_names, moe_expert_gate_names,
24 moe_expert_up_names, moe_router_names, output_names, output_norm_names,
25 per_layer_model_projection_names, per_layer_projection_norm_names, pre_feedforward_norm_names,
26 vision_proj_names,
27};
28use crate::tokenizer::{decode_placeholders, encode_naive, BundleTokenizer};
29use aria_kernel::{
30 attention_causal_with_scale, attention_with_scale, gated_delta_step, geglu, gelu_pytorch_tanh,
31 hdm_linear, kv_sliding_view, linear_cpu, moe_topk_route, resolve_compute, rms_norm,
32 rms_norm_gemma, rope_half, rope_half_partial, rope_half_proportional, short_conv_step,
33 silu_vec, softplus, swiglu, ComputeBackend, ComputePref, CudaContext, EngineError,
34 GatedDeltaStep,
35};
36use std::cell::RefCell;
37use std::collections::HashMap;
38use std::path::Path;
39use std::sync::Arc;
40use std::time::Instant;
41
42#[derive(Debug, Clone)]
43pub struct GenerateOpts {
44 pub max_tokens: usize,
45 pub temperature: f32,
46}
47
48impl Default for GenerateOpts {
49 fn default() -> Self {
50 Self {
51 max_tokens: 16,
52 temperature: 0.0,
53 }
54 }
55}
56
57#[derive(Debug, Clone)]
58pub struct Generation {
59 pub tokens: Vec<u32>,
60 pub text: String,
61}
62
63#[derive(Clone)]
64struct MatWeight {
65 data: Arc<Vec<f32>>,
66 hdm_seed: Option<i64>,
67}
68
69impl MatWeight {
70 fn from_loaded(w: LoadedWeight) -> Self {
71 Self {
72 data: Arc::new(w.data),
73 hdm_seed: w.hdm_seed,
74 }
75 }
76
77 fn concat_out(a: &Self, b: &Self) -> Self {
79 let mut data = Vec::with_capacity(a.data.len() + b.data.len());
80 data.extend_from_slice(&a.data);
81 data.extend_from_slice(&b.data);
82 Self {
83 data: Arc::new(data),
84 hdm_seed: None,
85 }
86 }
87}
88
89#[derive(Clone, Copy)]
90enum GemmAcct {
91 Attn,
92 Ffn,
93 LmHead,
94 Other,
95}
96
97#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
98enum AttnKind {
99 Sliding,
100 Full,
101}
102
103#[derive(Debug, Clone, Copy, PartialEq)]
106enum RopeMode {
107 Full,
108 Partial(f32),
109 Proportional(f32),
110}
111
112struct DecodeState {
114 k_caches: Vec<Vec<f32>>,
115 v_caches: Vec<Vec<f32>>,
116 last_kv_src: HashMap<AttnKind, usize>,
117 conv_states: Vec<Option<Vec<f32>>>,
118 delta_states: Vec<Option<Vec<f32>>>,
119 pos: usize,
121}
122
123#[derive(Clone)]
124struct AttnWeights {
125 wq: MatWeight,
126 wk: Option<MatWeight>,
128 wv: Option<MatWeight>,
129 wo: MatWeight,
130 q_norm: Option<Vec<f32>>,
131 k_norm: Option<Vec<f32>>,
132 v_norm: Option<Vec<f32>>,
133 kind: AttnKind,
134 q_gate: bool,
137}
138
139#[derive(Clone)]
140struct ConvWeights {
141 in_proj: MatWeight,
142 out_proj: MatWeight,
143 kernel: Vec<f32>,
145 kernel_size: usize,
146}
147
148#[derive(Clone)]
150struct DeltaWeights {
151 qkvz: MatWeight,
152 ba: MatWeight,
153 conv: Vec<f32>,
154 conv_k: usize,
155 out_proj: MatWeight,
156 out_norm: Vec<f32>,
158 a_log: Vec<f32>,
159 dt_bias: Vec<f32>,
160 n_k_heads: usize,
161 n_v_heads: usize,
162 head_k: usize,
163 head_v: usize,
164}
165
166#[derive(Clone)]
167enum LayerOp {
168 Attn(AttnWeights),
169 Conv(ConvWeights),
170 Linear(DeltaWeights),
171}
172
173#[derive(Clone)]
174struct ExpertWeights {
175 gate: MatWeight,
176 up: MatWeight,
177 down: MatWeight,
178}
179
180#[derive(Clone)]
181enum FfnWeights {
182 Dense {
183 gate: MatWeight,
184 up: MatWeight,
185 down: MatWeight,
186 },
187 MoE {
188 router: MatWeight,
189 experts: Vec<ExpertWeights>,
190 top_k: usize,
191 use_sigmoid: bool,
192 },
193}
194
195struct LayerPle {
196 gate: MatWeight,
197 proj: MatWeight,
198 post_norm: Vec<f32>,
199}
200
201struct PleModel {
202 embed: Arc<Vec<f32>>,
203 proj: MatWeight,
204 proj_norm: Vec<f32>,
205 d: usize,
206}
207
208struct LayerAltUp {
210 modality_router: MatWeight,
211 router_norm: Vec<f32>,
212 prediction_coefs: MatWeight,
213 correction_coefs: MatWeight,
214 correct_output_scale: Vec<f32>,
215}
216
217struct LayerLaurel {
219 left: MatWeight,
220 right: MatWeight,
221 post_norm: Vec<f32>,
222 rank: usize,
223}
224
225struct LayerWeights {
226 attn_norm: Vec<f32>,
227 ffn_norm: Vec<f32>,
228 post_attn_norm: Option<Vec<f32>>,
229 post_ffn_norm: Option<Vec<f32>>,
230 ple: Option<LayerPle>,
231 altup: Option<LayerAltUp>,
232 laurel: Option<LayerLaurel>,
233 activation_sparsity: f32,
235 layer_scalar: f32,
237 op: LayerOp,
238 ffn: FfnWeights,
239}
240
241struct ModelWeights {
242 emb: MatWeight,
243 layers: Vec<LayerWeights>,
244 output_norm: Vec<f32>,
245 output: MatWeight,
246 vision: Option<MatWeight>,
247 action: Option<MatWeight>,
248 ple: Option<PleModel>,
249 altup_projections: Vec<MatWeight>,
251 altup_unembed: Vec<MatWeight>,
252}
253
254pub struct Session {
255 family: Family,
256 bundle: Bundle,
257 weights: ModelWeights,
258 conf: crate::bundle::ModelConfig,
259 use_gemma_norm: bool,
260 use_gemma4: bool,
261 use_gemma3n: bool,
262 use_geglu: bool,
263 embed_scale: f32,
264 final_logit_softcap: Option<f32>,
266 tokenizer: Option<BundleTokenizer>,
267 decode: Option<DecodeState>,
269 compute: ComputeBackend,
270 compute_label: String,
271 cuda: Option<CudaContext>,
272 profile_on: bool,
273 last_profile: Option<EngineProfile>,
274 gen_acc: RefCell<GenerateProfile>,
275}
276
277impl std::fmt::Debug for Session {
278 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
279 f.debug_struct("Session")
280 .field("family", &self.family)
281 .field("model", &self.conf.hidden_size)
282 .finish()
283 }
284}
285
286pub struct SessionBuilder {
287 path: Option<std::path::PathBuf>,
288 family_path: String,
289 family_explicit: bool,
290 compute: ComputePref,
291 profile: bool,
292}
293
294impl SessionBuilder {
295 pub fn new() -> Self {
296 Self {
297 path: None,
298 family_path: "gemma/gemma-4-e2b-it".into(),
299 family_explicit: false,
300 compute: ComputePref::Auto,
301 profile: false,
302 }
303 }
304
305 pub fn model(mut self, path: impl AsRef<Path>) -> Self {
306 self.path = Some(path.as_ref().to_path_buf());
307 self
308 }
309
310 pub fn family(mut self, path: impl Into<String>) -> Self {
311 self.family_path = path.into();
312 self.family_explicit = true;
313 self
314 }
315
316 pub fn compute(mut self, pref: ComputePref) -> Self {
317 self.compute = pref;
318 self
319 }
320
321 pub fn profile(mut self, on: bool) -> Self {
322 self.profile = on;
323 self
324 }
325
326 pub fn build(self) -> Result<Session, EngineError> {
327 let path = self
328 .path
329 .ok_or_else(|| EngineError::InvalidParam("model path required".into()))?;
330 let family_path = if self.family_explicit {
334 self.family_path
335 } else {
336 path.to_str()
337 .and_then(infer_family_path)
338 .map(str::to_string)
339 .unwrap_or(self.family_path)
340 };
341 let family = require_runnable(&family_path)?;
342 let _hook = graph_hook(family.arch);
343 let (compute, compute_label) = resolve_compute(self.compute)?;
344 load_profile_begin(self.profile);
345 let t_mmap = Instant::now();
346 let bundle = load_bundle(&path)?;
347 load_profile_set_mmap(elapsed_ms(t_mmap));
348 let mut conf = bundle.model.clone();
349 conf.rope_theta = effective_rope_theta(family.path(), conf.rope_theta);
350 fill_gemma4_architecture_defaults(&mut conf, family.path());
353 fill_gemma3_architecture_defaults(&mut conf, family.path());
354 fill_gemma3n_architecture_defaults(&mut conf, family.path());
355 fill_qwen35_architecture_defaults(&mut conf, family.path());
356 require_gemma4_config(&conf, family.path())?;
357 reject_unsupported_geometry(&conf, family)?;
358 let tokenizer = BundleTokenizer::try_load(&path)?;
359 let t_mat = Instant::now();
360 let weights = materialize_with_config(&bundle, family, &conf)?;
363 require_gemma4_ple(&weights, &conf, family.path())?;
364 require_gemma3n_altup(&weights, &conf, family.path())?;
365 load_profile_set_materialize(elapsed_ms(t_mat));
366 let mut cuda = None;
367 if compute == ComputeBackend::Cuda {
368 let t_up = Instant::now();
369 let ctx = CudaContext::new()?;
370 upload_weights(&ctx, &weights)?;
371 load_profile_set_cuda_upload(elapsed_ms(t_up));
372 cuda = Some(ctx);
373 }
374 let act = conf
375 .hidden_act
376 .as_deref()
377 .unwrap_or("")
378 .to_ascii_lowercase();
379 let use_gemma4 = family.path().contains("gemma-4");
380 let use_gemma3n = is_gemma3n(family.path());
381 let use_gemma_norm = (family.path().contains("gemma") && !use_gemma4 && !use_gemma3n)
384 || family.path().contains("qwen3.5");
385 let use_geglu = act.contains("gelu") || family.path().contains("gemma");
386 let embed_scale = if family.path().contains("gemma") {
389 (conf.hidden_size as f32).sqrt()
390 } else {
391 1.0
392 };
393 let final_logit_softcap = if use_gemma4 || use_gemma3n {
394 Some(30.0)
395 } else {
396 None
397 };
398 let load = load_profile_take();
399 let last_profile = self.profile.then(|| EngineProfile {
400 compute: compute_label.clone(),
401 load,
402 generate: None,
403 ci_fail: false,
404 });
405 Ok(Session {
406 family,
407 bundle,
408 weights,
409 conf,
410 use_gemma_norm,
411 use_gemma4,
412 use_gemma3n,
413 use_geglu,
414 embed_scale,
415 final_logit_softcap,
416 tokenizer,
417 decode: None,
418 compute,
419 compute_label,
420 cuda,
421 profile_on: self.profile,
422 last_profile,
423 gen_acc: RefCell::new(GenerateProfile::default()),
424 })
425 }
426}
427
428impl Default for SessionBuilder {
429 fn default() -> Self {
430 Self::new()
431 }
432}
433
434fn reject_unsupported_geometry(
435 conf: &crate::bundle::ModelConfig,
436 family: Family,
437) -> Result<(), EngineError> {
438 let path = family.path();
439 if path.contains("qwen3.5") || path.contains("bonsai") {
441 let has_linear = conf
442 .layer_types
443 .as_ref()
444 .map(|t| {
445 t.iter().any(|s| {
446 let s = s.to_ascii_lowercase();
447 s.contains("linear_attention") || s.contains("delta")
448 })
449 })
450 .unwrap_or(false);
451 if !has_linear {
452 return Err(EngineError::Unsupported(format!(
453 "{path}: requires model.layer_types with Gated DeltaNet / linear_attention \
454 (dense-only bundles are unsupported until DeltaNet lands)"
455 )));
456 }
457 }
458 if family.is_moe() && conf.num_experts.unwrap_or(0) == 0 {
460 return Err(EngineError::Unsupported(format!(
461 "{path}: MoE family requires model.num_experts > 0 in bundle config"
462 )));
463 }
464 Ok(())
465}
466
467fn layer_type_str(conf: &crate::bundle::ModelConfig, layer: usize) -> String {
468 conf.layer_types
469 .as_ref()
470 .and_then(|t| t.get(layer))
471 .map(|s| s.to_ascii_lowercase())
472 .unwrap_or_else(|| "full_attention".into())
473}
474
475fn layer_is_conv(conf: &crate::bundle::ModelConfig, layer: usize) -> bool {
476 layer_type_str(conf, layer).contains("conv")
477}
478
479fn layer_is_linear(conf: &crate::bundle::ModelConfig, layer: usize) -> bool {
480 let t = layer_type_str(conf, layer);
481 t.contains("linear_attention") || t.contains("delta")
482}
483
484fn attn_kind(conf: &crate::bundle::ModelConfig, layer: usize) -> AttnKind {
485 if layer_type_str(conf, layer).contains("sliding") {
486 AttnKind::Sliding
487 } else {
488 AttnKind::Full
489 }
490}
491
492fn is_kv_consumer(conf: &crate::bundle::ModelConfig, layer: usize) -> bool {
493 let n = conf.num_kv_shared_layers.unwrap_or(0);
494 n > 0 && layer >= conf.num_layers.saturating_sub(n)
495}
496
497fn default_gemma4_layer_types(n: usize) -> Vec<String> {
499 (0..n)
500 .map(|i| {
501 if (i + 1) % 5 == 0 {
502 "full_attention".into()
503 } else {
504 "sliding_attention".into()
505 }
506 })
507 .collect()
508}
509
510fn is_gemma3_text(path: &str) -> bool {
511 path.to_ascii_lowercase().contains("gemma-3-")
513}
514
515fn is_gemma3n(path: &str) -> bool {
516 path.to_ascii_lowercase().contains("gemma-3n")
517}
518
519const GEMMA3N_ALTUP_N: usize = 4;
520
521fn altup_router_input_scale(hidden: usize) -> f32 {
523 if hidden == 0 {
524 0.0
525 } else {
526 1.0 / hidden as f32
527 }
528}
529
530fn gemma3n_default_kv_shared(num_layers: usize) -> usize {
532 num_layers.saturating_sub(20.min(num_layers))
533}
534
535fn default_gemma3_layer_types(n: usize) -> Vec<String> {
537 (0..n)
538 .map(|i| {
539 if (i + 1) % 6 == 0 {
540 "full_attention".into()
541 } else {
542 "sliding_attention".into()
543 }
544 })
545 .collect()
546}
547
548fn fill_gemma3_architecture_defaults(conf: &mut crate::bundle::ModelConfig, family_path: &str) {
552 if !is_gemma3_text(family_path) {
553 return;
554 }
555 if conf.layer_types.as_ref().map(|t| t.len()) != Some(conf.num_layers) {
556 conf.layer_types = Some(default_gemma3_layer_types(conf.num_layers));
557 }
558 if conf.sliding_window.unwrap_or(0) == 0 {
559 conf.sliding_window = Some(512);
560 }
561 if conf
562 .hidden_act
563 .as_ref()
564 .map(|s| s.is_empty())
565 .unwrap_or(true)
566 {
567 conf.hidden_act = Some("gelu_pytorch_tanh".into());
568 }
569}
570
571fn fill_gemma3n_architecture_defaults(conf: &mut crate::bundle::ModelConfig, family_path: &str) {
575 if !is_gemma3n(family_path) {
576 return;
577 }
578 if conf.layer_types.as_ref().map(|t| t.len()) != Some(conf.num_layers) {
579 conf.layer_types = Some(default_gemma4_layer_types(conf.num_layers));
580 }
581 if conf.sliding_window.unwrap_or(0) == 0 {
582 conf.sliding_window = Some(512);
583 }
584 if conf
585 .hidden_act
586 .as_ref()
587 .map(|s| s.is_empty())
588 .unwrap_or(true)
589 {
590 conf.hidden_act = Some("gelu_pytorch_tanh".into());
591 }
592 if conf.hidden_size >= 1024 {
593 if conf.head_dim.unwrap_or(0) == 0 {
594 conf.head_dim = Some(256);
595 }
596 if conf.num_kv_shared_layers.is_none() {
597 conf.num_kv_shared_layers = Some(gemma3n_default_kv_shared(conf.num_layers));
598 }
599 }
600}
601
602fn fill_qwen35_architecture_defaults(conf: &mut crate::bundle::ModelConfig, family_path: &str) {
606 if !family_path.contains("qwen3.5") {
607 return;
608 }
609 if conf.partial_rotary_factor.is_none() {
610 conf.partial_rotary_factor = Some(0.25);
611 }
612}
613
614fn fill_gemma4_architecture_defaults(conf: &mut crate::bundle::ModelConfig, family_path: &str) {
617 if !family_path.contains("gemma-4") {
618 return;
619 }
620 if conf.layer_types.as_ref().map(|t| t.len()) != Some(conf.num_layers) {
621 conf.layer_types = Some(default_gemma4_layer_types(conf.num_layers));
622 }
623 if conf.sliding_window.unwrap_or(0) == 0 {
624 conf.sliding_window = Some(512);
625 }
626 if conf.partial_rotary_factor.is_none() {
627 conf.partial_rotary_factor = Some(0.25);
628 }
629 if conf.hidden_size >= 1024 {
631 if conf.head_dim.unwrap_or(0) == 0 {
632 conf.head_dim = Some(256);
633 }
634 if conf.global_head_dim.unwrap_or(0) == 0 {
635 conf.global_head_dim = Some(512);
636 }
637 if conf.num_kv_shared_layers.is_none() {
638 conf.num_kv_shared_layers = Some(20);
639 }
640 } else {
641 if conf.head_dim.unwrap_or(0) == 0 && conf.num_attention_heads > 0 {
642 conf.head_dim = Some(conf.hidden_size / conf.num_attention_heads);
643 }
644 if conf.global_head_dim.unwrap_or(0) == 0 {
645 conf.global_head_dim = conf.head_dim;
646 }
647 }
648}
649
650fn require_gemma4_config(
652 conf: &crate::bundle::ModelConfig,
653 family_path: &str,
654) -> Result<(), EngineError> {
655 if !family_path.contains("gemma-4") {
656 return Ok(());
657 }
658 let missing = |field: &str| {
659 EngineError::Unsupported(format!(
660 "{family_path}: model.{field} required after Gemma-4 architecture fill \
661 (re-quantize with current model config_from_hf)"
662 ))
663 };
664 match &conf.layer_types {
665 None => return Err(missing("layer_types")),
666 Some(t) if t.len() != conf.num_layers => {
667 return Err(EngineError::Unsupported(format!(
668 "{family_path}: model.layer_types length {} != num_layers {}",
669 t.len(),
670 conf.num_layers
671 )));
672 }
673 Some(_) => {}
674 }
675 if conf.sliding_window.unwrap_or(0) == 0 {
676 return Err(missing("sliding_window"));
677 }
678 match conf.partial_rotary_factor {
679 Some(f) if f > 0.0 && f <= 1.0 => {}
680 _ => return Err(missing("partial_rotary_factor")),
681 }
682 if conf.head_dim.unwrap_or(0) == 0 {
683 return Err(missing("head_dim"));
684 }
685 if conf.global_head_dim.unwrap_or(0) == 0 {
686 return Err(missing("global_head_dim"));
687 }
688 Ok(())
689}
690
691fn gemma4_requires_ple(family_path: &str, hidden: usize) -> bool {
692 (family_path.contains("gemma-4") || is_gemma3n(family_path)) && hidden >= 1024
693}
694
695fn gemma3n_requires_altup(family_path: &str, hidden: usize) -> bool {
696 is_gemma3n(family_path) && hidden >= 1024
697}
698
699fn require_gemma3n_altup(
700 weights: &ModelWeights,
701 conf: &crate::bundle::ModelConfig,
702 family_path: &str,
703) -> Result<(), EngineError> {
704 if !gemma3n_requires_altup(family_path, conf.hidden_size) {
705 return Ok(());
706 }
707 let n_extra = GEMMA3N_ALTUP_N - 1;
708 if weights.altup_projections.len() != n_extra || weights.altup_unembed.len() != n_extra {
709 return Err(EngineError::Format(format!(
710 "{family_path}: Gemma-3n AltUp projections required \
711 (altup_projections + altup_unembed_projections); refusing silent no-op"
712 )));
713 }
714 for (i, layer) in weights.layers.iter().enumerate() {
715 if layer.altup.is_none() || layer.laurel.is_none() {
716 return Err(EngineError::Format(format!(
717 "{family_path}: Gemma-3n layer {i} missing AltUp/Laurel weights"
718 )));
719 }
720 }
721 Ok(())
722}
723
724fn require_gemma4_ple(
725 weights: &ModelWeights,
726 conf: &crate::bundle::ModelConfig,
727 family_path: &str,
728) -> Result<(), EngineError> {
729 if !gemma4_requires_ple(family_path, conf.hidden_size) {
730 return Ok(());
731 }
732 if weights.ple.is_none() {
733 return Err(EngineError::Format(format!(
734 "{family_path}: codebook PLE required for Gemma-4 / Gemma-3n E2B/E4B \
735 (embed_tokens_per_layer + per_layer_model_projection + \
736 per_layer_projection_norm); refusing silent no-op"
737 )));
738 }
739 Ok(())
740}
741
742fn resolve_attn_kind(
744 conf: &crate::bundle::ModelConfig,
745 layer: usize,
746 q_dim: usize,
747 n_heads: usize,
748) -> AttnKind {
749 if n_heads > 0 && q_dim.is_multiple_of(n_heads) {
750 let head_from_q = q_dim / n_heads;
751 if let (Some(g), Some(h)) = (
752 conf.global_head_dim.filter(|d| *d > 0),
753 conf.head_dim.filter(|d| *d > 0),
754 ) {
755 if g != h {
756 if head_from_q == g {
757 return AttnKind::Full;
758 }
759 if head_from_q == h {
760 return AttnKind::Sliding;
761 }
762 }
763 }
764 }
765 attn_kind(conf, layer)
766}
767
768fn attn_q_geometry(
770 wq_len: usize,
771 wo_len: usize,
772 hidden: usize,
773) -> Result<(usize, usize, bool), EngineError> {
774 if hidden == 0 || !wq_len.is_multiple_of(hidden) {
775 return Err(EngineError::ShapeMismatch(format!(
776 "attn q proj weight not divisible by hidden_size (len={wq_len} hidden={hidden})"
777 )));
778 }
779 if !wo_len.is_multiple_of(hidden) {
780 return Err(EngineError::ShapeMismatch(format!(
781 "attn output proj weight not divisible by hidden_size (len={wo_len} hidden={hidden})"
782 )));
783 }
784 let q_out = wq_len / hidden;
785 let wo_in = wo_len / hidden;
786 if q_out == wo_in {
787 Ok((q_out, q_out, false))
788 } else if q_out == 2 * wo_in {
789 Ok((q_out, wo_in, true))
790 } else {
791 Err(EngineError::ShapeMismatch(format!(
792 "attn output proj weight shape mismatch (wo_len={wo_len} hidden={hidden} q_out={q_out})"
793 )))
794 }
795}
796
797fn split_interleaved_q_gate(
799 mixed: &[f32],
800 seq: usize,
801 n_heads: usize,
802 head_dim: usize,
803) -> Result<(Vec<f32>, Vec<f32>), EngineError> {
804 let q_dim = n_heads.saturating_mul(head_dim);
805 let packed = q_dim.saturating_mul(2);
806 if packed == 0 || mixed.len() != seq * packed {
807 return Err(EngineError::ShapeMismatch(format!(
808 "gated q_proj out {} != seq*2*q_dim {}*{packed}",
809 mixed.len(),
810 seq
811 )));
812 }
813 let mut q = vec![0.0f32; seq * q_dim];
814 let mut gate = vec![0.0f32; seq * q_dim];
815 for t in 0..seq {
816 for h in 0..n_heads {
817 let src = t * packed + h * (2 * head_dim);
818 let dst = t * q_dim + h * head_dim;
819 q[dst..dst + head_dim].copy_from_slice(&mixed[src..src + head_dim]);
820 gate[dst..dst + head_dim].copy_from_slice(&mixed[src + head_dim..src + 2 * head_dim]);
821 }
822 }
823 Ok((q, gate))
824}
825
826fn apply_sigmoid_gate(x: &mut [f32], gate: &[f32]) -> Result<(), EngineError> {
827 if x.len() != gate.len() {
828 return Err(EngineError::ShapeMismatch(format!(
829 "attn output gate len {} != attn out {}",
830 gate.len(),
831 x.len()
832 )));
833 }
834 for (v, g) in x.iter_mut().zip(gate) {
835 *v *= 1.0 / (1.0 + (-*g).exp());
836 }
837 Ok(())
838}
839
840fn materialize_with_config(
841 b: &Bundle,
842 family: Family,
843 conf: &crate::bundle::ModelConfig,
844) -> Result<ModelWeights, EngineError> {
845 fn any_mat(b: &Bundle, names: &[String]) -> Result<MatWeight, EngineError> {
846 let refs: Vec<&str> = names.iter().map(String::as_str).collect();
847 Ok(MatWeight::from_loaded(b.weight_loaded_any(&refs)?))
848 }
849 fn try_mat(b: &Bundle, names: &[String]) -> Result<Option<MatWeight>, EngineError> {
850 match any_mat(b, names) {
851 Ok(w) => Ok(Some(w)),
852 Err(EngineError::Format(_)) => Ok(None),
853 Err(e) => Err(e),
854 }
855 }
856 fn any_vec(b: &Bundle, names: &[String]) -> Result<Vec<f32>, EngineError> {
857 Ok((*any_mat(b, names)?.data).clone())
858 }
859 fn optional_vec(b: &Bundle, names: &[String]) -> Option<Vec<f32>> {
860 any_vec(b, names).ok()
861 }
862
863 let m = conf;
864 let hidden = m.hidden_size;
865 let n_heads = m.num_attention_heads;
866 let n_experts = m.num_experts.unwrap_or(0);
867 let top_k = m.num_experts_per_tok.unwrap_or(1).max(1);
868 let use_sigmoid_router = n_experts > 0 && m.layer_types.is_some();
870
871 let mut layers = Vec::with_capacity(m.num_layers);
872 let mut prev_wk: Option<MatWeight> = None;
873 let mut prev_wv: Option<MatWeight> = None;
874 for layer in 0..m.num_layers {
875 let attn_norm = any_vec(b, &attn_norm_names(layer))?;
876 let pre_ff = optional_vec(b, &pre_feedforward_norm_names(layer));
877 let post_attn_norm = if pre_ff.is_some() {
878 optional_vec(b, &attn_post_norm_names(layer))
879 } else {
880 None
881 };
882 let post_ffn_norm = optional_vec(b, &ffn_post_norm_names(layer));
883 let ffn_norm = if let Some(v) = pre_ff {
884 v
885 } else {
886 any_vec(b, &ffn_norm_names(layer))?
887 };
888
889 let op = if layer_is_conv(m, layer) {
890 let in_proj = any_mat(b, &conv_in_proj_names(layer))?;
891 let out_proj = any_mat(b, &conv_out_proj_names(layer))?;
892 let kw = any_mat(b, &conv_kernel_names(layer))?;
893 let kernel_size = m.conv_l_cache.unwrap_or(3).max(1);
894 if kw.data.len() % hidden != 0 {
895 return Err(EngineError::ShapeMismatch(format!(
896 "layer {layer} conv kernel len {} not divisible by hidden {hidden}",
897 kw.data.len()
898 )));
899 }
900 let inferred_k = kw.data.len() / hidden;
901 let kernel_size = if inferred_k > 0 {
902 inferred_k
903 } else {
904 kernel_size
905 };
906 if kw.data.len() != hidden * kernel_size {
908 return Err(EngineError::ShapeMismatch(format!(
909 "layer {layer} conv kernel len {} != hidden*kernel {hidden}*{kernel_size}",
910 kw.data.len()
911 )));
912 }
913 let kernel = (*kw.data).clone();
914 if in_proj.data.len() != 3 * hidden * hidden {
915 return Err(EngineError::ShapeMismatch(format!(
916 "layer {layer} conv in_proj len {} != 3*hidden*hidden",
917 in_proj.data.len()
918 )));
919 }
920 if out_proj.data.len() != hidden * hidden {
921 return Err(EngineError::ShapeMismatch(format!(
922 "layer {layer} conv out_proj len {} != hidden*hidden",
923 out_proj.data.len()
924 )));
925 }
926 LayerOp::Conv(ConvWeights {
927 in_proj,
928 out_proj,
929 kernel,
930 kernel_size,
931 })
932 } else if layer_is_linear(m, layer) {
933 let qkvz = if let Some(w) = try_mat(b, &linear_in_proj_qkvz_names(layer))? {
934 w
935 } else {
936 let qkv = any_mat(b, &linear_in_proj_qkv_names(layer))?;
937 let z = any_mat(b, &linear_in_proj_z_names(layer))?;
938 MatWeight::concat_out(&qkv, &z)
939 };
940 let ba = if let Some(w) = try_mat(b, &linear_in_proj_ba_names(layer))? {
941 w
942 } else {
943 let proj_b = any_mat(b, &linear_in_proj_b_names(layer))?;
944 let proj_a = any_mat(b, &linear_in_proj_a_names(layer))?;
945 MatWeight::concat_out(&proj_b, &proj_a)
946 };
947 let conv_w = any_mat(b, &linear_conv1d_names(layer))?;
948 let out_proj = any_mat(b, &linear_out_proj_names(layer))?;
949 let a_log = any_vec(b, &linear_a_log_names(layer))?;
950 let dt_bias = any_vec(b, &linear_dt_bias_names(layer))?;
951 let n_v_heads = a_log.len();
952 if n_v_heads == 0 || dt_bias.len() != n_v_heads {
953 return Err(EngineError::ShapeMismatch(format!(
954 "layer {layer} A_log/dt_bias head mismatch"
955 )));
956 }
957 if ba.data.len() % hidden != 0 {
958 return Err(EngineError::ShapeMismatch(
959 "linear in_proj_ba not divisible by hidden".into(),
960 ));
961 }
962 if ba.data.len() / hidden != 2 * n_v_heads {
963 return Err(EngineError::ShapeMismatch(format!(
964 "layer {layer} in_proj_ba out {} != 2*n_v_heads {}",
965 ba.data.len() / hidden,
966 2 * n_v_heads
967 )));
968 }
969 if qkvz.data.len() % hidden != 0 {
970 return Err(EngineError::ShapeMismatch(
971 "linear in_proj_qkvz not divisible by hidden".into(),
972 ));
973 }
974 let qkvz_out = qkvz.data.len() / hidden;
975 if !qkvz_out.is_multiple_of(4) {
977 return Err(EngineError::ShapeMismatch(format!(
978 "layer {layer} qkvz out {qkvz_out} not divisible by 4"
979 )));
980 }
981 let key_dim = qkvz_out / 4;
982 let value_dim = key_dim;
983 let n_k_heads = n_v_heads;
984 if n_k_heads == 0
985 || !key_dim.is_multiple_of(n_k_heads)
986 || !value_dim.is_multiple_of(n_v_heads)
987 {
988 return Err(EngineError::ShapeMismatch(format!(
989 "layer {layer} cannot infer DeltaNet head dims"
990 )));
991 }
992 let head_k = key_dim / n_k_heads;
993 let head_v = value_dim / n_v_heads;
994 let conv_dim = key_dim * 2 + value_dim;
995 if conv_w.data.len() % conv_dim != 0 {
996 return Err(EngineError::ShapeMismatch(format!(
997 "layer {layer} conv1d len {} not divisible by conv_dim {conv_dim}",
998 conv_w.data.len()
999 )));
1000 }
1001 let conv_k = conv_w.data.len() / conv_dim;
1002 if out_proj.data.len() != hidden * value_dim {
1003 return Err(EngineError::ShapeMismatch(format!(
1004 "layer {layer} linear out_proj len {} != hidden*value_dim",
1005 out_proj.data.len()
1006 )));
1007 }
1008 let out_norm = optional_vec(b, &linear_out_norm_names(layer))
1009 .unwrap_or_else(|| vec![1.0f32; head_v]);
1010 if out_norm.len() != head_v {
1011 return Err(EngineError::ShapeMismatch(format!(
1012 "layer {layer} linear_attn.norm len {} != head_v {head_v}",
1013 out_norm.len()
1014 )));
1015 }
1016 LayerOp::Linear(DeltaWeights {
1017 qkvz,
1018 ba,
1019 conv: (*conv_w.data).clone(),
1020 conv_k,
1021 out_proj,
1022 out_norm,
1023 a_log,
1024 dt_bias,
1025 n_k_heads,
1026 n_v_heads,
1027 head_k,
1028 head_v,
1029 })
1030 } else {
1031 let consumer = is_kv_consumer(m, layer);
1032 let (wk, wv) = if consumer {
1033 (None, None)
1034 } else {
1035 let wk = match any_mat(b, &attn_k_names(layer)) {
1036 Ok(w) => {
1037 prev_wk = Some(w.clone());
1038 Some(w)
1039 }
1040 Err(e) => Some(prev_wk.clone().ok_or_else(|| {
1041 EngineError::Format(format!(
1042 "missing k_proj for layer {layer} and no prior KV to share ({e})"
1043 ))
1044 })?),
1045 };
1046 let wv = match any_mat(b, &attn_v_names(layer)) {
1047 Ok(w) => {
1048 prev_wv = Some(w.clone());
1049 Some(w)
1050 }
1051 Err(e) => Some(prev_wv.clone().ok_or_else(|| {
1052 EngineError::Format(format!(
1053 "missing v_proj for layer {layer} and no prior KV to share ({e})"
1054 ))
1055 })?),
1056 };
1057 (wk, wv)
1058 };
1059 let wq = any_mat(b, &attn_q_names(layer))?;
1060 let wo = any_mat(b, &attn_o_names(layer))?;
1061 let (_q_out, q_dim, q_gate) = attn_q_geometry(wq.data.len(), wo.data.len(), hidden)?;
1062 LayerOp::Attn(AttnWeights {
1063 wq,
1064 wk,
1065 wv,
1066 wo,
1067 q_norm: optional_vec(b, &attn_q_norm_names(layer)),
1068 k_norm: optional_vec(b, &attn_k_norm_names(layer)),
1069 v_norm: optional_vec(b, &attn_v_norm_names(layer)),
1070 kind: resolve_attn_kind(m, layer, q_dim, n_heads),
1071 q_gate,
1072 })
1073 };
1074
1075 let ffn = if n_experts > 0 {
1076 match any_mat(b, &moe_router_names(layer)) {
1077 Ok(router) => {
1078 if router.data.len() != n_experts * hidden {
1079 return Err(EngineError::ShapeMismatch(format!(
1080 "layer {layer} MoE router len {} != num_experts*hidden {n_experts}*{hidden}",
1081 router.data.len()
1082 )));
1083 }
1084 let mut experts = Vec::with_capacity(n_experts);
1085 for e in 0..n_experts {
1086 experts.push(ExpertWeights {
1087 gate: any_mat(b, &moe_expert_gate_names(layer, e))?,
1088 up: any_mat(b, &moe_expert_up_names(layer, e))?,
1089 down: any_mat(b, &moe_expert_down_names(layer, e))?,
1090 });
1091 }
1092 FfnWeights::MoE {
1093 router,
1094 experts,
1095 top_k,
1096 use_sigmoid: use_sigmoid_router,
1097 }
1098 }
1099 Err(_) => {
1100 FfnWeights::Dense {
1102 gate: any_mat(b, &ffn_gate_names(layer))?,
1103 up: any_mat(b, &ffn_up_names(layer))?,
1104 down: any_mat(b, &ffn_down_names(layer))?,
1105 }
1106 }
1107 }
1108 } else {
1109 FfnWeights::Dense {
1110 gate: any_mat(b, &ffn_gate_names(layer))?,
1111 up: any_mat(b, &ffn_up_names(layer))?,
1112 down: any_mat(b, &ffn_down_names(layer))?,
1113 }
1114 };
1115
1116 let ple = match (
1117 any_mat(b, &layer_ple_gate_names(layer)),
1118 any_mat(b, &layer_ple_proj_names(layer)),
1119 optional_vec(b, &layer_ple_post_norm_names(layer)),
1120 ) {
1121 (Ok(gate), Ok(proj), Some(post_norm)) => Some(LayerPle {
1122 gate,
1123 proj,
1124 post_norm,
1125 }),
1126 _ => None,
1127 };
1128
1129 let altup = match (
1130 try_mat(b, &layer_altup_router_names(layer))?,
1131 optional_vec(b, &layer_altup_router_norm_names(layer)),
1132 try_mat(b, &layer_altup_prediction_coef_names(layer))?,
1133 try_mat(b, &layer_altup_correction_coef_names(layer))?,
1134 optional_vec(b, &layer_altup_correct_scale_names(layer)),
1135 ) {
1136 (
1137 Some(modality_router),
1138 Some(router_norm),
1139 Some(prediction_coefs),
1140 Some(correction_coefs),
1141 Some(correct_output_scale),
1142 ) => Some(LayerAltUp {
1143 modality_router,
1144 router_norm,
1145 prediction_coefs,
1146 correction_coefs,
1147 correct_output_scale,
1148 }),
1149 (None, None, None, None, None) => None,
1150 _ => {
1151 return Err(EngineError::Format(format!(
1152 "layer {layer}: partial Gemma-3n AltUp tensors (need router, norms, pred/corr coefs, correct_output_scale)"
1153 )));
1154 }
1155 };
1156
1157 let laurel = match (
1158 try_mat(b, &layer_laurel_left_names(layer))?,
1159 try_mat(b, &layer_laurel_right_names(layer))?,
1160 optional_vec(b, &layer_laurel_norm_names(layer)),
1161 ) {
1162 (Some(left), Some(right), Some(post_norm)) => {
1163 if hidden == 0 || left.data.len() % hidden != 0 {
1164 return Err(EngineError::ShapeMismatch(format!(
1165 "layer {layer} laurel left not divisible by hidden"
1166 )));
1167 }
1168 let rank = left.data.len() / hidden;
1169 if rank == 0 || right.data.len() != hidden * rank {
1170 return Err(EngineError::ShapeMismatch(format!(
1171 "layer {layer} laurel rank/shape mismatch"
1172 )));
1173 }
1174 Some(LayerLaurel {
1175 left,
1176 right,
1177 post_norm,
1178 rank,
1179 })
1180 }
1181 (None, None, None) => None,
1182 _ => {
1183 return Err(EngineError::Format(format!(
1184 "layer {layer}: partial Gemma-3n Laurel tensors"
1185 )));
1186 }
1187 };
1188
1189 let activation_sparsity = if is_gemma3n(family.path()) && m.num_layers >= 30 && layer < 10 {
1190 0.95
1191 } else {
1192 0.0
1193 };
1194
1195 let layer_scalar = optional_vec(b, &layer_scalar_names(layer))
1196 .and_then(|v| v.into_iter().find(|x| x.is_finite()))
1197 .unwrap_or(1.0);
1198
1199 layers.push(LayerWeights {
1200 attn_norm,
1201 ffn_norm,
1202 post_attn_norm,
1203 post_ffn_norm,
1204 ple,
1205 altup,
1206 laurel,
1207 activation_sparsity,
1208 layer_scalar,
1209 op,
1210 ffn,
1211 });
1212 }
1213 let emb_n = emb_names();
1214 let out_norm_n = output_norm_names();
1215 let out_n = output_names();
1216 let vis_n: Vec<String> = vision_proj_names()
1217 .iter()
1218 .map(|s| (*s).to_string())
1219 .collect();
1220 let act_n: Vec<String> = action_head_names()
1221 .iter()
1222 .map(|s| (*s).to_string())
1223 .collect();
1224 let emb = MatWeight::from_loaded(b.weight_loaded_any(&emb_n)?);
1225 let output = if m.tie_word_embeddings.unwrap_or(false)
1226 || family.path().contains("gemma-4")
1227 || is_gemma3n(family.path())
1228 || (family.path().contains("qwen3") && !family.path().contains("qwen3.5"))
1229 {
1230 emb.clone()
1233 } else {
1234 MatWeight::from_loaded(b.weight_loaded_any(&out_n)?)
1235 };
1236 let require_ple = gemma4_requires_ple(family.path(), hidden);
1237 let ple = {
1238 let embed_n = embed_per_layer_names();
1239 let proj_n: Vec<String> = per_layer_model_projection_names()
1240 .iter()
1241 .map(|s| (*s).to_string())
1242 .collect();
1243 let norm_n: Vec<String> = per_layer_projection_norm_names()
1244 .iter()
1245 .map(|s| (*s).to_string())
1246 .collect();
1247 let embed_res = b.weight_loaded_any(&embed_n);
1248 let proj_res = any_mat(b, &proj_n);
1249 let proj_norm = optional_vec(b, &norm_n);
1250 match (embed_res, proj_res, proj_norm) {
1251 (Ok(embed), Ok(proj), Some(proj_norm)) => {
1252 let d = proj_norm.len();
1253 if d == 0 {
1254 return Err(EngineError::ShapeMismatch(
1255 "PLE projection norm dim is 0".into(),
1256 ));
1257 }
1258 Some(PleModel {
1259 embed: Arc::new(embed.data),
1260 proj,
1261 proj_norm,
1262 d,
1263 })
1264 }
1265 (embed_res, proj_res, proj_norm) if require_ple => {
1266 let embed_s = match &embed_res {
1267 Ok(_) => "ok".to_string(),
1268 Err(e) => e.to_string(),
1269 };
1270 let proj_s = match &proj_res {
1271 Ok(_) => "ok".to_string(),
1272 Err(e) => e.to_string(),
1273 };
1274 let norm_s = if proj_norm.is_some() { "ok" } else { "missing" };
1275 return Err(EngineError::Format(format!(
1276 "{}: codebook PLE required (embed_tokens_per_layer={embed_s}, \
1277 per_layer_model_projection={proj_s}, per_layer_projection_norm={norm_s})",
1278 family.path()
1279 )));
1280 }
1281 _ => None,
1282 }
1283 };
1284 if let Some(ple) = &ple {
1285 let packed = m.num_layers.saturating_mul(ple.d);
1286 if packed == 0
1287 || !ple.embed.len().is_multiple_of(packed)
1288 || ple.proj.data.len() != packed * hidden
1289 {
1290 return Err(EngineError::ShapeMismatch(format!(
1291 "PLE shapes: embed {} proj {} expected packed={} hidden={hidden}",
1292 ple.embed.len(),
1293 ple.proj.data.len(),
1294 packed
1295 )));
1296 }
1297 for (i, layer) in layers.iter().enumerate() {
1298 let Some(lp) = &layer.ple else {
1299 return Err(EngineError::Format(format!(
1300 "PLE model tensors present but layer {i} missing gate/proj/norm"
1301 )));
1302 };
1303 if lp.gate.data.len() != ple.d * hidden || lp.proj.data.len() != hidden * ple.d {
1304 return Err(EngineError::ShapeMismatch(format!(
1305 "layer {i} PLE gate/proj shape mismatch (d={}, hidden={hidden})",
1306 ple.d
1307 )));
1308 }
1309 if lp.post_norm.len() != hidden {
1312 return Err(EngineError::ShapeMismatch(format!(
1313 "layer {i} PLE post_norm len {} != hidden {hidden}",
1314 lp.post_norm.len()
1315 )));
1316 }
1317 }
1318 }
1319 let n_extra = GEMMA3N_ALTUP_N - 1;
1320 let mut altup_projections = Vec::new();
1321 let mut altup_unembed = Vec::new();
1322 let mut altup_partial = false;
1323 for i in 0..n_extra {
1324 match try_mat(b, &altup_projection_names(i))? {
1325 Some(w) => altup_projections.push(w),
1326 None => altup_partial = true,
1327 }
1328 match try_mat(b, &altup_unembed_names(i))? {
1329 Some(w) => altup_unembed.push(w),
1330 None => altup_partial = true,
1331 }
1332 }
1333 if altup_partial {
1334 if !altup_projections.is_empty() || !altup_unembed.is_empty() {
1335 return Err(EngineError::Format(
1336 "partial Gemma-3n altup_projections / altup_unembed_projections".into(),
1337 ));
1338 }
1339 altup_projections.clear();
1340 altup_unembed.clear();
1341 } else {
1342 for w in altup_projections.iter().chain(altup_unembed.iter()) {
1343 if w.data.len() != hidden * hidden {
1344 return Err(EngineError::ShapeMismatch(format!(
1345 "Gemma-3n altup projection len {} != hidden² {hidden}",
1346 w.data.len()
1347 )));
1348 }
1349 }
1350 }
1351 Ok(ModelWeights {
1352 emb,
1353 layers,
1354 output_norm: b.weight_loaded_any(&out_norm_n)?.data,
1355 output,
1356 vision: any_mat(b, &vis_n).ok(),
1357 action: any_mat(b, &act_n).ok(),
1358 ple,
1359 altup_projections,
1360 altup_unembed,
1361 })
1362}
1363
1364fn upload_weights(ctx: &CudaContext, w: &ModelWeights) -> Result<(), EngineError> {
1365 ctx.upload(&w.emb.data)?;
1366 ctx.upload(&w.output.data)?;
1367 if let Some(v) = &w.vision {
1368 ctx.upload(&v.data)?;
1369 }
1370 if let Some(a) = &w.action {
1371 ctx.upload(&a.data)?;
1372 }
1373 if let Some(ple) = &w.ple {
1374 ctx.upload(&ple.embed)?;
1375 ctx.upload(&ple.proj.data)?;
1376 }
1377 for wproj in w.altup_projections.iter().chain(w.altup_unembed.iter()) {
1378 ctx.upload(&wproj.data)?;
1379 }
1380 for layer in &w.layers {
1381 match &layer.op {
1382 LayerOp::Attn(attn) => {
1383 ctx.upload(&attn.wq.data)?;
1384 if let Some(wk) = &attn.wk {
1385 ctx.upload(&wk.data)?;
1386 }
1387 if let Some(wv) = &attn.wv {
1388 ctx.upload(&wv.data)?;
1389 }
1390 ctx.upload(&attn.wo.data)?;
1391 }
1392 LayerOp::Conv(c) => {
1393 ctx.upload(&c.in_proj.data)?;
1394 ctx.upload(&c.out_proj.data)?;
1395 }
1396 LayerOp::Linear(d) => {
1397 ctx.upload(&d.qkvz.data)?;
1398 ctx.upload(&d.ba.data)?;
1399 ctx.upload(&d.out_proj.data)?;
1400 }
1401 }
1402 if let Some(ple) = &layer.ple {
1403 ctx.upload(&ple.gate.data)?;
1404 ctx.upload(&ple.proj.data)?;
1405 }
1406 if let Some(altup) = &layer.altup {
1407 ctx.upload(&altup.modality_router.data)?;
1408 ctx.upload(&altup.prediction_coefs.data)?;
1409 ctx.upload(&altup.correction_coefs.data)?;
1410 }
1411 if let Some(laurel) = &layer.laurel {
1412 ctx.upload(&laurel.left.data)?;
1413 ctx.upload(&laurel.right.data)?;
1414 }
1415 match &layer.ffn {
1416 FfnWeights::Dense { gate, up, down } => {
1417 ctx.upload(&gate.data)?;
1418 ctx.upload(&up.data)?;
1419 ctx.upload(&down.data)?;
1420 }
1421 FfnWeights::MoE {
1422 router, experts, ..
1423 } => {
1424 ctx.upload(&router.data)?;
1425 for e in experts {
1426 ctx.upload(&e.gate.data)?;
1427 ctx.upload(&e.up.data)?;
1428 ctx.upload(&e.down.data)?;
1429 }
1430 }
1431 }
1432 }
1433 Ok(())
1434}
1435
1436impl Session {
1437 pub fn family(&self) -> Family {
1438 self.family
1439 }
1440
1441 pub fn model_id(&self) -> &str {
1442 self.family.path()
1443 }
1444
1445 pub fn config(&self) -> &crate::bundle::ModelConfig {
1446 &self.conf
1447 }
1448
1449 pub fn bundle(&self) -> &Bundle {
1450 &self.bundle
1451 }
1452
1453 pub fn compute_label(&self) -> &str {
1454 &self.compute_label
1455 }
1456
1457 pub fn last_profile(&self) -> Option<&EngineProfile> {
1458 self.last_profile.as_ref()
1459 }
1460
1461 fn wmm(
1462 &self,
1463 w: &MatWeight,
1464 x: &[f32],
1465 out_f: usize,
1466 in_f: usize,
1467 acct: GemmAcct,
1468 ) -> Result<Vec<f32>, EngineError> {
1469 let t0 = Instant::now();
1470 let y = if let Some(seed) = w.hdm_seed {
1471 hdm_linear(x, &w.data, out_f, in_f, Some(seed))?
1472 } else if self.compute == ComputeBackend::Cuda {
1473 let ctx = self.cuda.as_ref().ok_or_else(|| {
1474 EngineError::Unsupported("compute=cuda but CudaContext missing".into())
1475 })?;
1476 ctx.linear(x, &w.data, out_f, in_f)?
1477 } else {
1478 linear_cpu(x, &w.data, out_f, in_f)?
1479 };
1480 if self.profile_on {
1481 let ms = elapsed_ms(t0);
1482 let mut g = self.gen_acc.borrow_mut();
1483 match acct {
1484 GemmAcct::Attn => g.gemm_attn_ms += ms,
1485 GemmAcct::Ffn => g.gemm_ffn_ms += ms,
1486 GemmAcct::LmHead => g.gemm_lm_head_ms += ms,
1487 GemmAcct::Other => {}
1488 }
1489 }
1490 Ok(y)
1491 }
1492
1493 fn can_batch_prefill(&self) -> bool {
1494 self.weights.layers.iter().all(|layer| {
1495 matches!(layer.op, LayerOp::Attn(_)) && matches!(layer.ffn, FfnWeights::Dense { .. })
1496 })
1497 }
1498
1499 pub fn generate(
1502 &mut self,
1503 prompt: &[u32],
1504 opts: &GenerateOpts,
1505 ) -> Result<Generation, EngineError> {
1506 if opts.max_tokens == 0 {
1507 return Err(EngineError::InvalidParam("max_tokens must be > 0".into()));
1508 }
1509 let mut tokens: Vec<u32> = prompt.to_vec();
1510 if tokens.is_empty() {
1511 tokens.push(1);
1512 }
1513 self.decode = Some(self.fresh_decode_state());
1514 *self.gen_acc.borrow_mut() = GenerateProfile::default();
1515 let result = (|| {
1516 let t_pre = Instant::now();
1517 let mut logits = if self.can_batch_prefill() && tokens.len() > 1 {
1518 self.forward_prompt(&tokens)?
1519 } else {
1520 let mut last = Vec::new();
1521 for &tok in &tokens {
1522 last = self.forward_step(tok)?;
1523 }
1524 last
1525 };
1526 if self.profile_on {
1527 self.gen_acc.borrow_mut().prefill_ms = elapsed_ms(t_pre);
1528 }
1529 let mut generated = Vec::new();
1530 let t_dec = Instant::now();
1531 for _ in 0..opts.max_tokens {
1532 let next = argmax(&logits);
1534 generated.push(next);
1535 tokens.push(next);
1536 if self.is_stop_id(next) {
1537 generated.pop();
1538 break;
1539 }
1540 logits = self.forward_step(next)?;
1541 }
1542 if self.profile_on {
1543 self.gen_acc.borrow_mut().decode_ms = elapsed_ms(t_dec);
1544 }
1545 let text = self.decode_tokens(&generated);
1546 Ok(Generation {
1547 tokens: generated,
1548 text,
1549 })
1550 })();
1551 if self.profile_on {
1552 let mut p = self.last_profile.take().unwrap_or(EngineProfile {
1553 compute: self.compute_label.clone(),
1554 load: load_profile_take(),
1555 generate: None,
1556 ci_fail: false,
1557 });
1558 p.generate = Some(self.gen_acc.borrow().clone());
1559 self.last_profile = Some(p);
1560 }
1561 self.decode = None;
1562 result
1563 }
1564
1565 pub fn decode_tokens(&self, ids: &[u32]) -> String {
1568 match &self.tokenizer {
1569 Some(tok) => {
1570 let raw = tok.decode_opts(ids, false);
1571 strip_assistant_visible(&raw)
1572 }
1573 None => decode_placeholders(ids),
1574 }
1575 }
1576
1577 pub fn encode_text(&self, text: &str) -> Vec<u32> {
1579 match &self.tokenizer {
1580 Some(tok) => match tok.encode(text) {
1581 Ok(ids) if !ids.is_empty() => ids,
1582 Ok(_) => encode_naive(text, self.conf.vocab_size as u32),
1583 Err(_) => encode_naive(text, self.conf.vocab_size as u32),
1584 },
1585 None => encode_naive(text, self.conf.vocab_size as u32),
1586 }
1587 }
1588
1589 pub fn encode_chat(&self, messages: &[ChatTurn]) -> Vec<u32> {
1591 let family = if self.family.path().contains("gemma-4") {
1594 self.family.path()
1595 } else {
1596 self.tokenizer
1597 .as_ref()
1598 .and_then(|t| t.chat_family_hint())
1599 .unwrap_or(self.family.path())
1600 };
1601 let prompt = apply_chat_template(family, messages);
1602 self.encode_text(&prompt)
1603 }
1604
1605 fn is_stop_id(&self, id: u32) -> bool {
1606 match &self.tokenizer {
1607 Some(t) => t.is_stop(id),
1608 None => id == 0,
1609 }
1610 }
1611
1612 pub fn arch(&self) -> ArchClass {
1613 self.family.arch
1614 }
1615
1616 pub fn graph_hook_name(&self) -> &'static str {
1617 graph_hook(self.family.arch)
1618 }
1619
1620 pub fn embed_text(&self, text: &str) -> Result<Vec<f32>, EngineError> {
1622 let toks = self.encode_text(text);
1623 let hidden = self.conf.hidden_size;
1624 let vocab = self.conf.vocab_size;
1625 let mut acc = vec![0.0f32; hidden];
1626 if toks.is_empty() {
1627 return Ok(acc);
1628 }
1629 for &tok in &toks {
1630 let tid = (tok as usize) % vocab;
1631 let row = &self.weights.emb.data[tid * hidden..(tid + 1) * hidden];
1632 for i in 0..hidden {
1633 acc[i] += row[i];
1634 }
1635 }
1636 let inv = 1.0 / toks.len() as f32;
1637 for v in &mut acc {
1638 *v *= inv;
1639 }
1640 Ok(acc)
1641 }
1642
1643 pub fn vision_prefix(
1645 &self,
1646 rgb: &[u8],
1647 height: usize,
1648 width: usize,
1649 ) -> Result<Vec<f32>, EngineError> {
1650 if !matches!(self.family.arch, ArchClass::VL | ArchClass::VLA) {
1651 return Err(EngineError::Unsupported(format!(
1652 "vision_prefix not available for arch {:?}",
1653 self.family.arch
1654 )));
1655 }
1656 let Some(proj) = &self.weights.vision else {
1657 return Err(EngineError::Unsupported(format!(
1658 "{}: no vision projector tensor in bundle",
1659 self.family.path()
1660 )));
1661 };
1662 let hidden = self.conf.hidden_size;
1663 if hidden == 0 || proj.data.len() % hidden != 0 {
1664 return Err(EngineError::ShapeMismatch(
1665 "vision projector not divisible by hidden_size".into(),
1666 ));
1667 }
1668 let in_f = proj.data.len() / hidden;
1669 let need = height
1670 .checked_mul(width)
1671 .and_then(|n| n.checked_mul(3))
1672 .ok_or_else(|| EngineError::InvalidParam("vision size overflow".into()))?;
1673 if rgb.len() < need {
1674 return Err(EngineError::ShapeMismatch(format!(
1675 "rgb len {} < {}x{}x3",
1676 rgb.len(),
1677 height,
1678 width
1679 )));
1680 }
1681 let mut feat = vec![0.0f32; in_f];
1682 let pixels = height * width;
1683 if in_f == 3 {
1684 let mut acc = [0.0f32; 3];
1685 for p in 0..pixels {
1686 acc[0] += rgb[p * 3] as f32 / 255.0;
1687 acc[1] += rgb[p * 3 + 1] as f32 / 255.0;
1688 acc[2] += rgb[p * 3 + 2] as f32 / 255.0;
1689 }
1690 let s = 1.0 / pixels.max(1) as f32;
1691 feat[0] = acc[0] * s;
1692 feat[1] = acc[1] * s;
1693 feat[2] = acc[2] * s;
1694 } else {
1695 for i in 0..in_f {
1696 feat[i] = rgb[i % need] as f32 / 255.0;
1697 }
1698 }
1699 self.wmm(proj, &feat, hidden, in_f, GemmAcct::Other)
1700 }
1701
1702 pub fn predict_action(&self, prompt: &str, action_dim: usize) -> Result<Vec<f32>, EngineError> {
1704 if self.family.arch != ArchClass::VLA {
1705 return Err(EngineError::Unsupported(format!(
1706 "predict_action requires VLA, got {:?}",
1707 self.family.arch
1708 )));
1709 }
1710 if action_dim == 0 {
1711 return Err(EngineError::InvalidParam("action_dim must be > 0".into()));
1712 }
1713 let Some(head) = &self.weights.action else {
1714 return Err(EngineError::Unsupported(format!(
1715 "{}: no action head tensor in bundle",
1716 self.family.path()
1717 )));
1718 };
1719 let h = self.embed_text(prompt)?;
1720 let hidden = self.conf.hidden_size;
1721 if head.data.len() % hidden != 0 {
1722 return Err(EngineError::ShapeMismatch(
1723 "action head not divisible by hidden_size".into(),
1724 ));
1725 }
1726 let out_f = head.data.len() / hidden;
1727 if out_f != action_dim {
1728 return Err(EngineError::ShapeMismatch(format!(
1729 "action head out {out_f} != requested {action_dim}"
1730 )));
1731 }
1732 self.wmm(head, &h, out_f, hidden, GemmAcct::Other)
1733 }
1734
1735 pub fn transcribe_pcm16le(&self, pcm: &[u8]) -> Result<String, EngineError> {
1737 asr_transcribe_pcm16le(pcm, self.conf.vocab_size as u32)
1738 }
1739
1740 fn norm(&self, x: &[f32], weight: &[f32]) -> Result<Vec<f32>, EngineError> {
1741 if self.use_gemma_norm {
1742 rms_norm_gemma(x, weight, 1e-6)
1743 } else {
1744 rms_norm(x, weight, 1e-6)
1745 }
1746 }
1747
1748 fn add_normed_residual(
1749 &self,
1750 x: &mut [f32],
1751 y: &[f32],
1752 post_norm: Option<&[f32]>,
1753 ) -> Result<(), EngineError> {
1754 if y.len() != x.len() {
1755 return Err(EngineError::ShapeMismatch(
1756 "residual length mismatch".into(),
1757 ));
1758 }
1759 if let Some(w) = post_norm {
1760 let yn = self.norm(y, w)?;
1761 for (a, b) in x.iter_mut().zip(yn.iter()) {
1762 *a += *b;
1763 }
1764 } else {
1765 for (a, b) in x.iter_mut().zip(y.iter()) {
1766 *a += *b;
1767 }
1768 }
1769 Ok(())
1770 }
1771
1772 fn attn_scale(&self, head_dim: usize) -> f32 {
1773 if self.use_gemma4 || self.use_gemma3n {
1774 1.0
1775 } else {
1776 1.0 / (head_dim as f32).sqrt()
1777 }
1778 }
1779
1780 fn attn_window(&self, kind: AttnKind) -> Option<usize> {
1784 if kind != AttnKind::Sliding {
1785 return None;
1786 }
1787 self.conf.sliding_window.filter(|w| *w > 0)
1788 }
1789
1790 fn layer_rope_params(&self, kind: AttnKind) -> (f32, RopeMode) {
1791 if self.use_gemma4 {
1792 match kind {
1793 AttnKind::Sliding => (10_000.0, RopeMode::Full),
1794 AttnKind::Full => {
1795 let factor = self.conf.partial_rotary_factor.unwrap_or(1.0);
1797 let mode = if factor > 0.0 && factor < 1.0 {
1798 RopeMode::Proportional(factor)
1799 } else {
1800 RopeMode::Full
1801 };
1802 (1_000_000.0, mode)
1803 }
1804 }
1805 } else if is_gemma3_text(self.family.path()) || self.use_gemma3n {
1806 match kind {
1809 AttnKind::Sliding => (10_000.0, RopeMode::Full),
1810 AttnKind::Full => {
1811 let theta = if (self.conf.rope_theta - 10_000.0).abs() < 0.5 {
1812 1_000_000.0
1813 } else {
1814 self.conf.rope_theta
1815 };
1816 (theta, RopeMode::Full)
1817 }
1818 }
1819 } else if self.family.path().contains("qwen3.5") {
1820 let factor = self.conf.partial_rotary_factor.unwrap_or(0.25);
1823 let mode = if factor > 0.0 && factor < 1.0 {
1824 RopeMode::Partial(factor)
1825 } else {
1826 RopeMode::Full
1827 };
1828 (self.conf.rope_theta, mode)
1829 } else {
1830 (self.conf.rope_theta, RopeMode::Full)
1831 }
1832 }
1833
1834 fn layer_head_dim(
1835 &self,
1836 kind: AttnKind,
1837 q_dim: usize,
1838 n_heads: usize,
1839 ) -> Result<usize, EngineError> {
1840 if n_heads == 0 || !q_dim.is_multiple_of(n_heads) {
1841 return Err(EngineError::ShapeMismatch(
1842 "q_dim not divisible by num_attention_heads".into(),
1843 ));
1844 }
1845 let configured = match kind {
1846 AttnKind::Full => self.conf.global_head_dim.or(self.conf.head_dim),
1847 AttnKind::Sliding => self.conf.head_dim,
1848 };
1849 Ok(configured
1850 .filter(|d| *d > 0 && q_dim == n_heads * *d)
1851 .unwrap_or(q_dim / n_heads))
1852 }
1853
1854 fn apply_rope(
1855 x: &mut [f32],
1856 head_dim: usize,
1857 pos: usize,
1858 theta: f32,
1859 mode: RopeMode,
1860 ) -> Result<(), EngineError> {
1861 match mode {
1862 RopeMode::Full => rope_half(x, head_dim, pos, theta),
1863 RopeMode::Proportional(factor) => {
1864 rope_half_proportional(x, head_dim, factor, pos, theta)
1867 }
1868 RopeMode::Partial(factor) => {
1869 let rotary_dim = (factor * head_dim as f32) as usize & !1;
1872 if rotary_dim < 2 || rotary_dim >= head_dim {
1873 rope_half(x, head_dim, pos, theta)
1874 } else {
1875 rope_half_partial(x, head_dim, rotary_dim, pos, theta)
1876 }
1877 }
1878 }
1879 }
1880
1881 fn apply_v_norm(
1882 &self,
1883 v: Vec<f32>,
1884 v_norm: Option<&[f32]>,
1885 head_dim: usize,
1886 ) -> Result<Vec<f32>, EngineError> {
1887 if let Some(vn) = v_norm {
1888 if vn.len() != head_dim {
1889 return Err(EngineError::ShapeMismatch(format!(
1890 "v_norm len {} != head_dim {head_dim}",
1891 vn.len()
1892 )));
1893 }
1894 self.norm(&v, vn)
1895 } else if self.use_gemma4 || self.use_gemma3n {
1896 let ones = vec![1.0f32; head_dim];
1897 rms_norm(&v, &ones, 1e-6)
1898 } else {
1899 Ok(v)
1900 }
1901 }
1902
1903 fn has_gemma3n_graph(&self) -> bool {
1904 self.use_gemma3n
1905 && self.weights.altup_projections.len() == GEMMA3N_ALTUP_N - 1
1906 && self.weights.altup_unembed.len() == GEMMA3N_ALTUP_N - 1
1907 }
1908
1909 fn match_token_magnitude(
1910 src: &[f32],
1911 target: &[f32],
1912 hidden: usize,
1913 ) -> Result<Vec<f32>, EngineError> {
1914 if src.len() != target.len() || hidden == 0 || !src.len().is_multiple_of(hidden) {
1915 return Err(EngineError::ShapeMismatch(
1916 "Gemma-3n magnitude match length mismatch".into(),
1917 ));
1918 }
1919 let seq = src.len() / hidden;
1920 let mut out = src.to_vec();
1921 for t in 0..seq {
1922 let tb = t * hidden;
1923 let mut t_ms = 0.0f32;
1924 let mut s_ms = 0.0f32;
1925 for i in 0..hidden {
1926 t_ms += target[tb + i] * target[tb + i];
1927 s_ms += src[tb + i] * src[tb + i];
1928 }
1929 let target_mag = (t_ms / hidden as f32).sqrt();
1930 let src_mag = (s_ms / hidden as f32).max(1e-5).sqrt();
1931 let scale = target_mag / src_mag;
1932 for i in 0..hidden {
1933 out[tb + i] *= scale;
1934 }
1935 }
1936 Ok(out)
1937 }
1938
1939 fn gemma3n_expand_streams(
1940 &self,
1941 x0: &[f32],
1942 hidden: usize,
1943 ) -> Result<Vec<Vec<f32>>, EngineError> {
1944 let mut streams = vec![x0.to_vec()];
1945 for proj in &self.weights.altup_projections {
1946 let p = self.wmm(proj, x0, hidden, hidden, GemmAcct::Other)?;
1947 streams.push(Self::match_token_magnitude(&p, x0, hidden)?);
1948 }
1949 Ok(streams)
1950 }
1951
1952 fn gemma3n_unembed_streams(
1953 &self,
1954 streams: &[Vec<f32>],
1955 hidden: usize,
1956 ) -> Result<Vec<f32>, EngineError> {
1957 if streams.len() != GEMMA3N_ALTUP_N {
1958 return Err(EngineError::Format(
1959 "Gemma-3n unembed expects 4 AltUp streams".into(),
1960 ));
1961 }
1962 let mut acc = streams[0].clone();
1963 for (i, proj) in self.weights.altup_unembed.iter().enumerate() {
1964 let u = self.wmm(proj, &streams[i + 1], hidden, hidden, GemmAcct::Other)?;
1965 let matched = Self::match_token_magnitude(&u, &streams[0], hidden)?;
1966 for (a, b) in acc.iter_mut().zip(matched.iter()) {
1967 *a += *b;
1968 }
1969 }
1970 let n = GEMMA3N_ALTUP_N as f32;
1971 for v in &mut acc {
1972 *v /= n;
1973 }
1974 Ok(acc)
1975 }
1976
1977 fn altup_modalities(
1978 &self,
1979 altup: &LayerAltUp,
1980 x: &[f32],
1981 hidden: usize,
1982 ) -> Result<Vec<f32>, EngineError> {
1983 if altup.router_norm.len() != hidden {
1984 return Err(EngineError::ShapeMismatch(
1985 "Gemma-3n altup.router_norm dim mismatch".into(),
1986 ));
1987 }
1988 let mut xn = self.norm(x, &altup.router_norm)?;
1989 let scale = altup_router_input_scale(hidden);
1990 for v in &mut xn {
1991 *v *= scale;
1992 }
1993 let mut routed = self.wmm(
1994 &altup.modality_router,
1995 &xn,
1996 GEMMA3N_ALTUP_N,
1997 hidden,
1998 GemmAcct::Other,
1999 )?;
2000 for v in &mut routed {
2001 *v = v.tanh();
2002 }
2003 Ok(routed)
2004 }
2005
2006 fn altup_predict(
2007 &self,
2008 layer: &LayerWeights,
2009 streams: &[Vec<f32>],
2010 hidden: usize,
2011 ) -> Result<Vec<Vec<f32>>, EngineError> {
2012 let altup = layer
2013 .altup
2014 .as_ref()
2015 .ok_or_else(|| EngineError::Format("Gemma-3n layer missing AltUp".into()))?;
2016 let n = GEMMA3N_ALTUP_N;
2017 if streams.len() != n {
2018 return Err(EngineError::Format(
2019 "Gemma-3n predict expects 4 streams".into(),
2020 ));
2021 }
2022 let seq = streams[0].len() / hidden;
2023 let routed = self.altup_modalities(altup, &streams[0], hidden)?;
2024 let coefs = self.wmm(&altup.prediction_coefs, &routed, n * n, n, GemmAcct::Other)?;
2025 let mut preds = vec![vec![0.0f32; seq * hidden]; n];
2026 for t in 0..seq {
2027 for out_s in 0..n {
2028 for h in 0..hidden {
2029 let mut acc = streams[out_s][t * hidden + h];
2030 for in_s in 0..n {
2031 acc += streams[in_s][t * hidden + h] * coefs[t * n * n + out_s * n + in_s];
2032 }
2033 preds[out_s][t * hidden + h] = acc;
2034 }
2035 }
2036 }
2037 Ok(preds)
2038 }
2039
2040 fn altup_correct(
2041 &self,
2042 layer: &LayerWeights,
2043 predictions: &[Vec<f32>],
2044 activated: &[f32],
2045 hidden: usize,
2046 ) -> Result<Vec<Vec<f32>>, EngineError> {
2047 let altup = layer
2048 .altup
2049 .as_ref()
2050 .ok_or_else(|| EngineError::Format("Gemma-3n layer missing AltUp".into()))?;
2051 let n = GEMMA3N_ALTUP_N;
2052 let seq = activated.len() / hidden;
2053 let routed = self.altup_modalities(altup, activated, hidden)?;
2054 let mut coefs = self.wmm(&altup.correction_coefs, &routed, n, n, GemmAcct::Other)?;
2055 for v in &mut coefs {
2056 *v += 1.0;
2057 }
2058 let mut out = Vec::with_capacity(n);
2059 for (si, pred) in predictions.iter().enumerate() {
2060 let mut row = pred.clone();
2061 for t in 0..seq {
2062 let c = coefs[t * n + si];
2063 for h in 0..hidden {
2064 let idx = t * hidden + h;
2065 let innov = activated[idx] - predictions[0][idx];
2066 row[idx] += innov * c;
2067 }
2068 }
2069 out.push(row);
2070 }
2071 Ok(out)
2072 }
2073
2074 fn apply_laurel(
2075 &self,
2076 layer: &LayerWeights,
2077 xn: &[f32],
2078 hidden: usize,
2079 ) -> Result<Vec<f32>, EngineError> {
2080 let laurel = layer
2081 .laurel
2082 .as_ref()
2083 .ok_or_else(|| EngineError::Format("Gemma-3n layer missing Laurel".into()))?;
2084 let left = self.wmm(&laurel.left, xn, laurel.rank, hidden, GemmAcct::Ffn)?;
2085 let right = self.wmm(&laurel.right, &left, hidden, laurel.rank, GemmAcct::Ffn)?;
2086 let nrm = self.norm(&right, &laurel.post_norm)?;
2087 let mut out = xn.to_vec();
2088 for (a, b) in out.iter_mut().zip(nrm.iter()) {
2089 *a += *b;
2090 }
2091 Ok(out)
2092 }
2093
2094 fn scale_altup_active(
2095 &self,
2096 layer: &LayerWeights,
2097 x: &mut [f32],
2098 hidden: usize,
2099 ) -> Result<(), EngineError> {
2100 let Some(altup) = &layer.altup else {
2101 return Ok(());
2102 };
2103 if altup.correct_output_scale.len() == 1 {
2104 let s = altup.correct_output_scale[0];
2105 for v in x.iter_mut() {
2106 *v *= s;
2107 }
2108 return Ok(());
2109 }
2110 if altup.correct_output_scale.len() != hidden {
2111 return Err(EngineError::ShapeMismatch(format!(
2112 "altup.correct_output_scale len {} != hidden {hidden}",
2113 altup.correct_output_scale.len()
2114 )));
2115 }
2116 let seq = x.len() / hidden;
2117 for t in 0..seq {
2118 for h in 0..hidden {
2119 x[t * hidden + h] *= altup.correct_output_scale[h];
2120 }
2121 }
2122 Ok(())
2123 }
2124
2125 fn gemma3n_after_attn(
2127 &self,
2128 streams: &mut [Vec<f32>],
2129 predictions: &[Vec<f32>],
2130 laurel: &[f32],
2131 ao: &[f32],
2132 li: usize,
2133 ple_tok: Option<&[f32]>,
2134 ) -> Result<(), EngineError> {
2135 let layer = &self.weights.layers[li];
2136 let hidden = self.conf.hidden_size;
2137 let mut active = predictions[0].clone();
2138 self.add_normed_residual(&mut active, ao, layer.post_attn_norm.as_deref())?;
2139 let inv_sqrt2 = std::f32::consts::FRAC_1_SQRT_2;
2140 let mut mix = vec![0.0f32; active.len()];
2141 for i in 0..active.len() {
2142 mix[i] = (active[i] + laurel[i]) * inv_sqrt2;
2143 }
2144 let xn2 = self.norm(&mix, &layer.ffn_norm)?;
2145 let down = self.apply_ffn(layer, &xn2, hidden)?;
2146 self.add_normed_residual(&mut mix, &down, layer.post_ffn_norm.as_deref())?;
2147 let corrected = self.altup_correct(layer, predictions, &mix, hidden)?;
2148 for (dst, src) in streams.iter_mut().zip(corrected) {
2149 *dst = src;
2150 }
2151 let mut first = streams[0].clone();
2152 self.scale_altup_active(layer, &mut first, hidden)?;
2153 if let Some(delta) = self.ple_delta(&first, layer, li, ple_tok, hidden)? {
2154 for s in streams.iter_mut().skip(1) {
2155 for (a, b) in s.iter_mut().zip(delta.iter()) {
2156 *a += *b;
2157 }
2158 }
2159 }
2160 Ok(())
2161 }
2162
2163 fn compute_ple_inputs(
2164 &self,
2165 toks: &[u32],
2166 embeds: &[f32],
2167 ) -> Result<Option<Vec<f32>>, EngineError> {
2168 let Some(ple) = &self.weights.ple else {
2169 return Ok(None);
2170 };
2171 let hidden = self.conf.hidden_size;
2172 let n_layers = self.weights.layers.len();
2173 let d = ple.d;
2174 let packed = n_layers * d;
2175 let seq = toks.len();
2176 if seq == 0 || embeds.len() != seq * hidden {
2177 return Err(EngineError::ShapeMismatch(
2178 "PLE embed sequence length mismatch".into(),
2179 ));
2180 }
2181 let scale_lookup = (d as f32).sqrt();
2182 let ple_vocab = ple.embed.len() / packed;
2183 if ple_vocab == 0 {
2184 return Err(EngineError::ShapeMismatch("PLE embed vocab is 0".into()));
2185 }
2186 let mut lookup = vec![0.0f32; seq * packed];
2187 for (t, &tok) in toks.iter().enumerate() {
2188 let tid = (tok as usize) % ple_vocab;
2189 let row = &ple.embed[tid * packed..(tid + 1) * packed];
2190 for i in 0..packed {
2191 lookup[t * packed + i] = row[i] * scale_lookup;
2192 }
2193 }
2194 let proj_scale = (hidden as f32).sqrt().recip();
2195 let mut proj = self.wmm(&ple.proj, embeds, packed, hidden, GemmAcct::Other)?;
2196 for v in &mut proj {
2197 *v *= proj_scale;
2198 }
2199 proj = rms_norm(&proj, &ple.proj_norm, 1e-6)?;
2200 let inv_sqrt2 = std::f32::consts::FRAC_1_SQRT_2;
2201 for i in 0..proj.len() {
2202 proj[i] = (proj[i] + lookup[i]) * inv_sqrt2;
2203 }
2204 Ok(Some(proj))
2205 }
2206
2207 fn ple_delta(
2208 &self,
2209 x: &[f32],
2210 layer: &LayerWeights,
2211 li: usize,
2212 ple_tok: Option<&[f32]>,
2213 hidden: usize,
2214 ) -> Result<Option<Vec<f32>>, EngineError> {
2215 let (Some(ple), Some(ple_tok)) = (&layer.ple, ple_tok) else {
2216 return Ok(None);
2217 };
2218 let d = self
2219 .weights
2220 .ple
2221 .as_ref()
2222 .map(|p| p.d)
2223 .ok_or_else(|| EngineError::Format("layer PLE without model PLE".into()))?;
2224 let n_layers = self.weights.layers.len();
2225 let seq = x.len() / hidden;
2226 let gate_out = self.wmm(&ple.gate, x, d, hidden, GemmAcct::Ffn)?;
2227 let mut gated = vec![0.0f32; seq * d];
2228 for t in 0..seq {
2229 for i in 0..d {
2230 let g = gelu_pytorch_tanh(gate_out[t * d + i]);
2231 let p = ple_tok[t * n_layers * d + li * d + i];
2232 gated[t * d + i] = g * p;
2233 }
2234 }
2235 let proj = self.wmm(&ple.proj, &gated, hidden, d, GemmAcct::Ffn)?;
2236 Ok(Some(self.norm(&proj, &ple.post_norm)?))
2237 }
2238
2239 fn apply_ple(
2240 &self,
2241 x: &mut [f32],
2242 layer: &LayerWeights,
2243 li: usize,
2244 ple_tok: Option<&[f32]>,
2245 hidden: usize,
2246 ) -> Result<(), EngineError> {
2247 if let Some(delta) = self.ple_delta(x, layer, li, ple_tok, hidden)? {
2248 for (a, b) in x.iter_mut().zip(delta.iter()) {
2249 *a += *b;
2250 }
2251 }
2252 Ok(())
2253 }
2254
2255 fn apply_layer_scalar(x: &mut [f32], scale: f32) {
2256 if (scale - 1.0).abs() < 1e-8 {
2257 return;
2258 }
2259 for v in x {
2260 *v *= scale;
2261 }
2262 }
2263
2264 fn apply_ffn(
2265 &self,
2266 layer: &LayerWeights,
2267 xn2: &[f32],
2268 hidden: usize,
2269 ) -> Result<Vec<f32>, EngineError> {
2270 match &layer.ffn {
2271 FfnWeights::Dense { gate, up, down } => {
2272 if gate.data.len() % hidden != 0 {
2273 return Err(EngineError::ShapeMismatch(
2274 "dense gate len not divisible by hidden".into(),
2275 ));
2276 }
2277 let inter = gate.data.len() / hidden;
2278 if inter == 0
2279 || up.data.len() != inter * hidden
2280 || down.data.len() != hidden * inter
2281 {
2282 return Err(EngineError::ShapeMismatch(
2283 "dense FFN weight shape mismatch".into(),
2284 ));
2285 }
2286 let mut g = self.wmm(gate, xn2, inter, hidden, GemmAcct::Ffn)?;
2287 if layer.activation_sparsity > 0.0 {
2288 g = gaussian_topk(&g, inter, layer.activation_sparsity)?;
2289 }
2290 let u = self.wmm(up, xn2, inter, hidden, GemmAcct::Ffn)?;
2291 let h = if self.use_geglu {
2292 geglu(&g, &u)?
2293 } else {
2294 swiglu(&g, &u)?
2295 };
2296 self.wmm(down, &h, hidden, inter, GemmAcct::Ffn)
2297 }
2298 FfnWeights::MoE {
2299 router,
2300 experts,
2301 top_k,
2302 use_sigmoid,
2303 } => {
2304 let n_exp = experts.len();
2305 let logits = self.wmm(router, xn2, n_exp, hidden, GemmAcct::Ffn)?;
2306 let (ids, weights) = moe_topk_route(&logits, *top_k, *use_sigmoid)?;
2307 let mut acc = vec![0.0f32; hidden];
2308 for (ei, &w) in ids.iter().zip(weights.iter()) {
2309 let ex = &experts[*ei];
2310 if ex.gate.data.len() % hidden != 0 {
2311 return Err(EngineError::ShapeMismatch(
2312 "expert gate len not divisible by hidden".into(),
2313 ));
2314 }
2315 let inter = ex.gate.data.len() / hidden;
2316 let g = self.wmm(&ex.gate, xn2, inter, hidden, GemmAcct::Ffn)?;
2317 let u = self.wmm(&ex.up, xn2, inter, hidden, GemmAcct::Ffn)?;
2318 let h = swiglu(&g, &u)?;
2319 let down = self.wmm(&ex.down, &h, hidden, inter, GemmAcct::Ffn)?;
2320 for i in 0..hidden {
2321 acc[i] += w * down[i];
2322 }
2323 }
2324 Ok(acc)
2325 }
2326 }
2327 }
2328
2329 fn fresh_decode_state(&self) -> DecodeState {
2330 let hidden = self.conf.hidden_size;
2331 DecodeState {
2332 k_caches: (0..self.conf.num_layers).map(|_| Vec::new()).collect(),
2333 v_caches: (0..self.conf.num_layers).map(|_| Vec::new()).collect(),
2334 last_kv_src: HashMap::new(),
2335 conv_states: self
2336 .weights
2337 .layers
2338 .iter()
2339 .map(|layer| match &layer.op {
2340 LayerOp::Conv(c) => {
2341 let hist = c.kernel_size.saturating_sub(1);
2342 Some(vec![0.0f32; hidden * hist])
2343 }
2344 LayerOp::Linear(d) => {
2345 let conv_dim = d.n_k_heads * d.head_k * 2 + d.n_v_heads * d.head_v;
2346 let hist = d.conv_k.saturating_sub(1);
2347 Some(vec![0.0f32; conv_dim * hist])
2348 }
2349 LayerOp::Attn(_) => None,
2350 })
2351 .collect(),
2352 delta_states: self
2353 .weights
2354 .layers
2355 .iter()
2356 .map(|layer| match &layer.op {
2357 LayerOp::Linear(d) => Some(vec![0.0f32; d.n_v_heads * d.head_k * d.head_v]),
2358 _ => None,
2359 })
2360 .collect(),
2361 pos: 0,
2362 }
2363 }
2364
2365 #[cfg(test)]
2367 fn forward(&self, tokens: &[u32]) -> Result<Vec<f32>, EngineError> {
2368 let mut state = self.fresh_decode_state();
2369 let mut logits = Vec::new();
2370 for &tok in tokens {
2371 logits = self.forward_step_with(&mut state, tok)?;
2372 }
2373 Ok(logits)
2374 }
2375
2376 fn forward_prompt(&mut self, toks: &[u32]) -> Result<Vec<f32>, EngineError> {
2377 let mut owned = self.decode.take().ok_or_else(|| {
2378 EngineError::InvalidParam("decode state missing; call generate/prefill first".into())
2379 })?;
2380 let logits = self.forward_prompt_with(&mut owned, toks);
2381 self.decode = Some(owned);
2382 logits
2383 }
2384
2385 fn apply_rope_seq(
2386 x: &mut [f32],
2387 seq: usize,
2388 tok_dim: usize,
2389 head_dim: usize,
2390 pos0: usize,
2391 theta: f32,
2392 mode: RopeMode,
2393 ) -> Result<(), EngineError> {
2394 if x.len() != seq * tok_dim {
2395 return Err(EngineError::ShapeMismatch(
2396 "rope seq buffer length mismatch".into(),
2397 ));
2398 }
2399 for t in 0..seq {
2400 Self::apply_rope(
2401 &mut x[t * tok_dim..(t + 1) * tok_dim],
2402 head_dim,
2403 pos0 + t,
2404 theta,
2405 mode,
2406 )?;
2407 }
2408 Ok(())
2409 }
2410
2411 fn forward_prompt_with(
2412 &self,
2413 state: &mut DecodeState,
2414 toks: &[u32],
2415 ) -> Result<Vec<f32>, EngineError> {
2416 if toks.is_empty() {
2417 return Err(EngineError::InvalidParam("empty prompt".into()));
2418 }
2419 let hidden = self.conf.hidden_size;
2420 let n_heads = self.conf.num_attention_heads;
2421 let n_kv = self.conf.num_kv_heads;
2422 let vocab = self.conf.vocab_size;
2423 let seq = toks.len();
2424 if hidden == 0 {
2425 return Err(EngineError::ShapeMismatch("hidden_size is 0".into()));
2426 }
2427 if self.weights.emb.data.len() < vocab.saturating_mul(hidden)
2428 || !self.weights.emb.data.len().is_multiple_of(hidden)
2429 {
2430 return Err(EngineError::ShapeMismatch(format!(
2431 "embedding length {} not compatible with vocab={vocab} hidden={hidden}",
2432 self.weights.emb.data.len()
2433 )));
2434 }
2435 let pos0 = state.pos;
2436 let mut x = vec![0.0f32; seq * hidden];
2437 for (t, &tok) in toks.iter().enumerate() {
2438 let tid = (tok as usize) % vocab;
2439 x[t * hidden..(t + 1) * hidden]
2440 .copy_from_slice(&self.weights.emb.data[tid * hidden..(tid + 1) * hidden]);
2441 }
2442 if self.embed_scale != 1.0 {
2443 for v in &mut x {
2444 *v *= self.embed_scale;
2445 }
2446 }
2447 let ple_tok = self.compute_ple_inputs(toks, &x)?;
2448 let mut gemma3n_streams = if self.has_gemma3n_graph() {
2449 Some(self.gemma3n_expand_streams(&x, hidden)?)
2450 } else {
2451 None
2452 };
2453
2454 for (li, layer) in self.weights.layers.iter().enumerate() {
2455 let mut gemma3n_preds = None;
2456 let mut gemma3n_laurel = None;
2457 let xn = if let Some(ref streams) = gemma3n_streams {
2458 let preds = self.altup_predict(layer, streams, hidden)?;
2459 let xn = self.norm(&preds[0], &layer.attn_norm)?;
2460 gemma3n_laurel = Some(self.apply_laurel(layer, &xn, hidden)?);
2461 gemma3n_preds = Some(preds);
2462 xn
2463 } else {
2464 self.norm(&x, &layer.attn_norm)?
2465 };
2466 match &layer.op {
2467 LayerOp::Attn(attn) => {
2468 if attn.wq.data.len() % hidden != 0 {
2469 return Err(EngineError::ShapeMismatch(
2470 "attn q proj weight not divisible by hidden_size".into(),
2471 ));
2472 }
2473 let q_out = attn.wq.data.len() / hidden;
2474 let q_dim = if attn.q_gate { q_out / 2 } else { q_out };
2475 let head_dim = self.layer_head_dim(attn.kind, q_dim, n_heads)?;
2476 if attn.wo.data.len() != hidden * q_dim {
2477 return Err(EngineError::ShapeMismatch(format!(
2478 "attn output proj weight shape mismatch (wo_len={} hidden={hidden} q_dim={q_dim})",
2479 attn.wo.data.len()
2480 )));
2481 }
2482 let mixed = self.wmm(&attn.wq, &xn, q_out, hidden, GemmAcct::Attn)?;
2483 let (mut q, gate) = if attn.q_gate {
2484 split_interleaved_q_gate(&mixed, seq, n_heads, head_dim)?
2485 } else {
2486 (mixed, Vec::new())
2487 };
2488 if let Some(qn) = &attn.q_norm {
2489 if qn.len() != head_dim {
2490 return Err(EngineError::ShapeMismatch(format!(
2491 "q_norm len {} != head_dim {head_dim}",
2492 qn.len()
2493 )));
2494 }
2495 q = self.norm(&q, qn)?;
2496 }
2497 let (theta, rope) = self.layer_rope_params(attn.kind);
2498 Self::apply_rope_seq(&mut q, seq, q_dim, head_dim, pos0, theta, rope)?;
2499
2500 let (k_src, v_src) = if let (Some(wk), Some(wv)) = (&attn.wk, &attn.wv) {
2501 if wk.data.len() % hidden != 0 || wv.data.len() % hidden != 0 {
2502 return Err(EngineError::ShapeMismatch(
2503 "attn kv proj weight not divisible by hidden_size".into(),
2504 ));
2505 }
2506 let k_dim = wk.data.len() / hidden;
2507 let v_dim = wv.data.len() / hidden;
2508 if k_dim != n_kv * head_dim || v_dim != n_kv * head_dim {
2509 return Err(EngineError::ShapeMismatch(format!(
2510 "kv dims {k_dim}/{v_dim} != n_kv*head_dim {}",
2511 n_kv * head_dim
2512 )));
2513 }
2514 let mut k = self.wmm(wk, &xn, k_dim, hidden, GemmAcct::Attn)?;
2515 let mut v = self.wmm(wv, &xn, v_dim, hidden, GemmAcct::Attn)?;
2516 if let Some(kn) = &attn.k_norm {
2517 if kn.len() != head_dim {
2518 return Err(EngineError::ShapeMismatch(format!(
2519 "k_norm len {} != head_dim {head_dim}",
2520 kn.len()
2521 )));
2522 }
2523 k = self.norm(&k, kn)?;
2524 }
2525 Self::apply_rope_seq(&mut k, seq, k_dim, head_dim, pos0, theta, rope)?;
2526 v = self.apply_v_norm(v, attn.v_norm.as_deref(), head_dim)?;
2527 state.k_caches[li] = k;
2528 state.v_caches[li] = v;
2529 state.last_kv_src.insert(attn.kind, li);
2530 (li, li)
2531 } else {
2532 let src = state.last_kv_src.get(&attn.kind).copied().ok_or_else(|| {
2533 EngineError::Format(format!(
2534 "KV-consumer layer {li} has no producer of kind {:?}",
2535 attn.kind
2536 ))
2537 })?;
2538 (src, src)
2539 };
2540 let attn_out = attention_causal_with_scale(
2541 &q,
2542 &state.k_caches[k_src],
2543 &state.v_caches[v_src],
2544 n_heads,
2545 n_kv,
2546 head_dim,
2547 self.attn_scale(head_dim),
2548 self.attn_window(attn.kind),
2549 )?;
2550 let mut attn_out = attn_out;
2551 if attn.q_gate {
2552 apply_sigmoid_gate(&mut attn_out, &gate)?;
2553 }
2554 let ao = self.wmm(&attn.wo, &attn_out, hidden, q_dim, GemmAcct::Attn)?;
2555 if let (Some(ref mut streams), Some(preds), Some(laurel)) =
2556 (gemma3n_streams.as_mut(), gemma3n_preds, gemma3n_laurel)
2557 {
2558 self.gemma3n_after_attn(
2559 streams,
2560 &preds,
2561 &laurel,
2562 &ao,
2563 li,
2564 ple_tok.as_deref(),
2565 )?;
2566 continue;
2567 }
2568 self.add_normed_residual(&mut x, &ao, layer.post_attn_norm.as_deref())?;
2569 }
2570 LayerOp::Conv(_) | LayerOp::Linear(_) => {
2571 return Err(EngineError::Unsupported(
2572 "batched prefill is only implemented for attention+dense FFN layers".into(),
2573 ));
2574 }
2575 }
2576 let xn2 = self.norm(&x, &layer.ffn_norm)?;
2577 let down = self.apply_ffn(layer, &xn2, hidden)?;
2578 self.add_normed_residual(&mut x, &down, layer.post_ffn_norm.as_deref())?;
2579 self.apply_ple(&mut x, layer, li, ple_tok.as_deref(), hidden)?;
2580 Self::apply_layer_scalar(&mut x, layer.layer_scalar);
2581 }
2582 if let Some(streams) = gemma3n_streams {
2583 x = self.gemma3n_unembed_streams(&streams, hidden)?;
2584 }
2585 state.pos = pos0 + seq;
2586 let last = &x[(seq - 1) * hidden..seq * hidden];
2587 let xn = self.norm(last, &self.weights.output_norm)?;
2588 if !self.weights.output.data.len().is_multiple_of(hidden) {
2589 return Err(EngineError::ShapeMismatch(format!(
2590 "lm_head len {} not divisible by hidden {hidden}",
2591 self.weights.output.data.len()
2592 )));
2593 }
2594 let out_rows = self.weights.output.data.len() / hidden;
2595 let logits = self.wmm(
2596 &self.weights.output,
2597 &xn,
2598 out_rows,
2599 hidden,
2600 GemmAcct::LmHead,
2601 )?;
2602 Ok(self.softcap_logits(logits))
2603 }
2604
2605 fn forward_step(&mut self, tok: u32) -> Result<Vec<f32>, EngineError> {
2606 let mut owned = self.decode.take().ok_or_else(|| {
2607 EngineError::InvalidParam("decode state missing; call generate/prefill first".into())
2608 })?;
2609 let logits = self.forward_step_with(&mut owned, tok);
2610 self.decode = Some(owned);
2611 logits
2612 }
2613
2614 fn forward_step_with(
2615 &self,
2616 state: &mut DecodeState,
2617 tok: u32,
2618 ) -> Result<Vec<f32>, EngineError> {
2619 let hidden = self.conf.hidden_size;
2620 let n_heads = self.conf.num_attention_heads;
2621 let n_kv = self.conf.num_kv_heads;
2622 let vocab = self.conf.vocab_size;
2623 if hidden == 0 {
2624 return Err(EngineError::ShapeMismatch("hidden_size is 0".into()));
2625 }
2626 if self.weights.emb.data.len() < vocab.saturating_mul(hidden)
2627 || !self.weights.emb.data.len().is_multiple_of(hidden)
2628 {
2629 return Err(EngineError::ShapeMismatch(format!(
2630 "embedding length {} not compatible with vocab={vocab} hidden={hidden}",
2631 self.weights.emb.data.len()
2632 )));
2633 }
2634 let pos = state.pos;
2635 let tid = (tok as usize) % vocab;
2636 let mut x = vec![0.0f32; hidden];
2637 x.copy_from_slice(&self.weights.emb.data[tid * hidden..(tid + 1) * hidden]);
2638 if self.embed_scale != 1.0 {
2639 for v in &mut x {
2640 *v *= self.embed_scale;
2641 }
2642 }
2643 let ple_tok = self.compute_ple_inputs(&[tok], &x)?;
2644 let mut gemma3n_streams = if self.has_gemma3n_graph() {
2645 Some(self.gemma3n_expand_streams(&x, hidden)?)
2646 } else {
2647 None
2648 };
2649
2650 for (li, layer) in self.weights.layers.iter().enumerate() {
2651 let mut gemma3n_preds = None;
2652 let mut gemma3n_laurel = None;
2653 let xn = if let Some(ref streams) = gemma3n_streams {
2654 let preds = self.altup_predict(layer, streams, hidden)?;
2655 let xn = self.norm(&preds[0], &layer.attn_norm)?;
2656 gemma3n_laurel = Some(self.apply_laurel(layer, &xn, hidden)?);
2657 gemma3n_preds = Some(preds);
2658 xn
2659 } else {
2660 self.norm(&x, &layer.attn_norm)?
2661 };
2662 match &layer.op {
2663 LayerOp::Attn(attn) => {
2664 if attn.wq.data.len() % hidden != 0 {
2665 return Err(EngineError::ShapeMismatch(
2666 "attn q proj weight not divisible by hidden_size".into(),
2667 ));
2668 }
2669 let q_out = attn.wq.data.len() / hidden;
2670 let q_dim = if attn.q_gate { q_out / 2 } else { q_out };
2671 let head_dim = self.layer_head_dim(attn.kind, q_dim, n_heads)?;
2672 if attn.wo.data.len() != hidden * q_dim {
2673 return Err(EngineError::ShapeMismatch(format!(
2674 "attn output proj weight shape mismatch (wo_len={} hidden={hidden} q_dim={q_dim})",
2675 attn.wo.data.len()
2676 )));
2677 }
2678 let mixed = self.wmm(&attn.wq, &xn, q_out, hidden, GemmAcct::Attn)?;
2679 let (mut q, gate) = if attn.q_gate {
2680 split_interleaved_q_gate(&mixed, 1, n_heads, head_dim)?
2681 } else {
2682 (mixed, Vec::new())
2683 };
2684 if let Some(qn) = &attn.q_norm {
2685 if qn.len() != head_dim {
2686 return Err(EngineError::ShapeMismatch(format!(
2687 "q_norm len {} != head_dim {head_dim}",
2688 qn.len()
2689 )));
2690 }
2691 q = self.norm(&q, qn)?;
2692 }
2693 let (theta, rope) = self.layer_rope_params(attn.kind);
2694 Self::apply_rope(&mut q, head_dim, pos, theta, rope)?;
2695
2696 let (k_src, v_src) = if let (Some(wk), Some(wv)) = (&attn.wk, &attn.wv) {
2697 if wk.data.len() % hidden != 0 || wv.data.len() % hidden != 0 {
2698 return Err(EngineError::ShapeMismatch(
2699 "attn kv proj weight not divisible by hidden_size".into(),
2700 ));
2701 }
2702 let k_dim = wk.data.len() / hidden;
2703 let v_dim = wv.data.len() / hidden;
2704 if k_dim != n_kv * head_dim || v_dim != n_kv * head_dim {
2705 return Err(EngineError::ShapeMismatch(format!(
2706 "kv dims {k_dim}/{v_dim} != n_kv*head_dim {}",
2707 n_kv * head_dim
2708 )));
2709 }
2710 let mut k = self.wmm(wk, &xn, k_dim, hidden, GemmAcct::Attn)?;
2711 let mut v = self.wmm(wv, &xn, v_dim, hidden, GemmAcct::Attn)?;
2712 if let Some(kn) = &attn.k_norm {
2713 if kn.len() != head_dim {
2714 return Err(EngineError::ShapeMismatch(format!(
2715 "k_norm len {} != head_dim {head_dim}",
2716 kn.len()
2717 )));
2718 }
2719 k = self.norm(&k, kn)?;
2720 }
2721 Self::apply_rope(&mut k, head_dim, pos, theta, rope)?;
2722 v = self.apply_v_norm(v, attn.v_norm.as_deref(), head_dim)?;
2723 state.k_caches[li].extend_from_slice(&k);
2724 state.v_caches[li].extend_from_slice(&v);
2725 state.last_kv_src.insert(attn.kind, li);
2726 (li, li)
2727 } else {
2728 let src = state.last_kv_src.get(&attn.kind).copied().ok_or_else(|| {
2729 EngineError::Format(format!(
2730 "KV-consumer layer {li} has no producer of kind {:?}",
2731 attn.kind
2732 ))
2733 })?;
2734 (src, src)
2735 };
2736 let kv_dim = n_kv * head_dim;
2737 let (k_view, v_view) = kv_sliding_view(
2738 &state.k_caches[k_src],
2739 &state.v_caches[v_src],
2740 kv_dim,
2741 self.attn_window(attn.kind),
2742 )?;
2743 let attn_out = attention_with_scale(
2744 &q,
2745 k_view,
2746 v_view,
2747 n_heads,
2748 n_kv,
2749 head_dim,
2750 self.attn_scale(head_dim),
2751 )?;
2752 let mut attn_out = attn_out;
2753 if attn.q_gate {
2754 apply_sigmoid_gate(&mut attn_out, &gate)?;
2755 }
2756 let ao = self.wmm(&attn.wo, &attn_out, hidden, q_dim, GemmAcct::Attn)?;
2757 if let (Some(ref mut streams), Some(preds), Some(laurel)) =
2758 (gemma3n_streams.as_mut(), gemma3n_preds, gemma3n_laurel)
2759 {
2760 self.gemma3n_after_attn(
2761 streams,
2762 &preds,
2763 &laurel,
2764 &ao,
2765 li,
2766 ple_tok.as_deref(),
2767 )?;
2768 continue;
2769 }
2770 self.add_normed_residual(&mut x, &ao, layer.post_attn_norm.as_deref())?;
2771 }
2772 LayerOp::Conv(conv) => {
2773 let bcx = self.wmm(&conv.in_proj, &xn, 3 * hidden, hidden, GemmAcct::Attn)?;
2774 let mut bx = vec![0.0f32; hidden];
2775 let mut c_gate = vec![0.0f32; hidden];
2776 for i in 0..hidden {
2777 let b = bcx[i];
2778 let c = bcx[hidden + i];
2779 let xx = bcx[2 * hidden + i];
2780 bx[i] = b * xx;
2781 c_gate[i] = c;
2782 }
2783 let cstate = state.conv_states[li]
2784 .as_mut()
2785 .ok_or_else(|| EngineError::ShapeMismatch("missing conv state".into()))?;
2786 let conv_y =
2787 short_conv_step(&bx, &conv.kernel, cstate, hidden, conv.kernel_size)?;
2788 let mut y = vec![0.0f32; hidden];
2789 for i in 0..hidden {
2790 y[i] = c_gate[i] * conv_y[i];
2791 }
2792 let ao = self.wmm(&conv.out_proj, &y, hidden, hidden, GemmAcct::Attn)?;
2793 self.add_normed_residual(&mut x, &ao, layer.post_attn_norm.as_deref())?;
2794 }
2795 LayerOp::Linear(dn) => {
2796 let key_dim = dn.n_k_heads * dn.head_k;
2797 let value_dim = dn.n_v_heads * dn.head_v;
2798 let qkvz_out = 2 * key_dim + 2 * value_dim;
2799 let mixed = self.wmm(&dn.qkvz, &xn, qkvz_out, hidden, GemmAcct::Attn)?;
2800 let mut q = mixed[0..key_dim].to_vec();
2801 let mut k = mixed[key_dim..2 * key_dim].to_vec();
2802 let mut v = mixed[2 * key_dim..2 * key_dim + value_dim].to_vec();
2803 let z = mixed[2 * key_dim + value_dim..].to_vec();
2804 let mut qkv = Vec::with_capacity(key_dim * 2 + value_dim);
2805 qkv.extend_from_slice(&q);
2806 qkv.extend_from_slice(&k);
2807 qkv.extend_from_slice(&v);
2808 let conv_dim = qkv.len();
2809 let cstate = state.conv_states[li].as_mut().ok_or_else(|| {
2810 EngineError::ShapeMismatch("missing delta conv state".into())
2811 })?;
2812 let mut mixed_c = short_conv_step(&qkv, &dn.conv, cstate, conv_dim, dn.conv_k)?;
2813 silu_vec(&mut mixed_c);
2814 q.copy_from_slice(&mixed_c[0..key_dim]);
2815 k.copy_from_slice(&mixed_c[key_dim..2 * key_dim]);
2816 v.copy_from_slice(&mixed_c[2 * key_dim..]);
2817 let ba = self.wmm(&dn.ba, &xn, 2 * dn.n_v_heads, hidden, GemmAcct::Attn)?;
2818 let mut beta = vec![0.0f32; dn.n_v_heads];
2819 let mut g = vec![0.0f32; dn.n_v_heads];
2820 for h in 0..dn.n_v_heads {
2821 beta[h] = 1.0 / (1.0 + (-ba[h]).exp());
2822 let alpha =
2823 -dn.a_log[h].exp() * softplus(ba[dn.n_v_heads + h] + dn.dt_bias[h]);
2824 g[h] = alpha.exp();
2825 }
2826 if dn.n_v_heads != dn.n_k_heads {
2827 return Err(EngineError::Unsupported(
2828 "DeltaNet GQA (n_v != n_k) not implemented".into(),
2829 ));
2830 }
2831 let s = state.delta_states[li].as_mut().ok_or_else(|| {
2832 EngineError::ShapeMismatch("missing delta recurrent state".into())
2833 })?;
2834 let mut core = gated_delta_step(GatedDeltaStep {
2835 q: &q,
2836 k: &k,
2837 v: &v,
2838 g: &g,
2839 beta: &beta,
2840 state: s,
2841 n_heads: dn.n_v_heads,
2842 dk: dn.head_k,
2843 dv: dn.head_v,
2844 })?;
2845 core = rms_norm(&core, &dn.out_norm, 1e-6)?;
2847 let mut z_act = z;
2848 silu_vec(&mut z_act);
2849 for i in 0..core.len() {
2850 core[i] *= z_act[i];
2851 }
2852 let ao = self.wmm(&dn.out_proj, &core, hidden, value_dim, GemmAcct::Attn)?;
2853 self.add_normed_residual(&mut x, &ao, layer.post_attn_norm.as_deref())?;
2854 }
2855 }
2856 let xn2 = self.norm(&x, &layer.ffn_norm)?;
2857 let down = self.apply_ffn(layer, &xn2, hidden)?;
2858 self.add_normed_residual(&mut x, &down, layer.post_ffn_norm.as_deref())?;
2859 self.apply_ple(&mut x, layer, li, ple_tok.as_deref(), hidden)?;
2860 Self::apply_layer_scalar(&mut x, layer.layer_scalar);
2861 }
2862 if let Some(streams) = gemma3n_streams {
2863 x = self.gemma3n_unembed_streams(&streams, hidden)?;
2864 }
2865 state.pos += 1;
2866 let xn = self.norm(&x, &self.weights.output_norm)?;
2867 if !self.weights.output.data.len().is_multiple_of(hidden) {
2868 return Err(EngineError::ShapeMismatch(format!(
2869 "lm_head len {} not divisible by hidden {hidden}",
2870 self.weights.output.data.len()
2871 )));
2872 }
2873 let out_rows = self.weights.output.data.len() / hidden;
2874 let logits = self.wmm(
2875 &self.weights.output,
2876 &xn,
2877 out_rows,
2878 hidden,
2879 GemmAcct::LmHead,
2880 )?;
2881 Ok(self.softcap_logits(logits))
2882 }
2883
2884 fn softcap_logits(&self, mut logits: Vec<f32>) -> Vec<f32> {
2885 if let Some(cap) = self.final_logit_softcap.filter(|c| *c > 0.0) {
2886 for x in &mut logits {
2887 *x = (*x / cap).tanh() * cap;
2888 }
2889 }
2890 logits
2891 }
2892}
2893
2894fn gaussian_topk(gate: &[f32], inter: usize, sparsity: f32) -> Result<Vec<f32>, EngineError> {
2896 if inter == 0 || !gate.len().is_multiple_of(inter) {
2897 return Err(EngineError::ShapeMismatch(
2898 "gaussian_topk: gate len not divisible by intermediate_size".into(),
2899 ));
2900 }
2901 let z = if (sparsity - 0.95).abs() < 0.02 {
2902 1.644_853_8
2903 } else {
2904 1.644_853_8 * (sparsity / 0.95).clamp(0.0, 4.0)
2906 };
2907 let seq = gate.len() / inter;
2908 let mut out = vec![0.0f32; gate.len()];
2909 let n = inter as f32;
2910 for t in 0..seq {
2911 let row = &gate[t * inter..(t + 1) * inter];
2912 let mean = row.iter().sum::<f32>() / n;
2913 let mut var = 0.0f32;
2914 for &v in row {
2915 let d = v - mean;
2916 var += d * d;
2917 }
2918 var /= n;
2919 let cutoff = mean + var.sqrt() * z;
2920 for i in 0..inter {
2921 out[t * inter + i] = (row[i] - cutoff).max(0.0);
2922 }
2923 }
2924 Ok(out)
2925}
2926
2927fn argmax(v: &[f32]) -> u32 {
2928 let mut best = 0usize;
2929 let mut best_v = f32::NEG_INFINITY;
2930 for (i, &x) in v.iter().enumerate() {
2931 if x > best_v {
2932 best_v = x;
2933 best = i;
2934 }
2935 }
2936 best as u32
2937}
2938
2939pub fn confidence_from_logits(logits: &[f32]) -> f32 {
2941 if logits.is_empty() {
2942 return 0.0;
2943 }
2944 let m = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
2945 let mut sum = 0.0f32;
2946 let mut maxp = 0.0f32;
2947 for &x in logits {
2948 let e = (x - m).exp();
2949 sum += e;
2950 if e > maxp {
2951 maxp = e;
2952 }
2953 }
2954 if sum > 0.0 {
2955 maxp / sum
2956 } else {
2957 0.0
2958 }
2959}
2960
2961#[allow(dead_code)]
2962pub fn cache_shapes_ok(cache: &HashMap<usize, Vec<f32>>, kv_dim: usize) -> bool {
2963 cache.values().all(|v| v.len().is_multiple_of(kv_dim))
2964}
2965
2966#[cfg(test)]
2967mod tests {
2968 use super::*;
2969 use crate::family::{arch_class_representatives, graph_hook, lookup_family, require_stage_b};
2970 use crate::fixture::write_tiny_q4_bundle;
2971 use aria_kernel::{resolve_compute, ComputePref};
2972 use serde_json::{json, Value};
2973
2974 #[test]
2975 fn gemma4_fills_hub_bundle_missing_geometry_fields() {
2976 let dir = tempfile::tempdir().unwrap();
2977 write_tiny_q4_bundle(dir.path()).unwrap();
2978 let cfg_path = dir.path().join("config.json");
2979 let raw = std::fs::read_to_string(&cfg_path).unwrap();
2980 let mut cfg: Value = serde_json::from_str(&raw).unwrap();
2981 let model = cfg["model"].as_object_mut().unwrap();
2983 for key in [
2984 "layer_types",
2985 "sliding_window",
2986 "partial_rotary_factor",
2987 "global_head_dim",
2988 "head_dim",
2989 ] {
2990 model.remove(key);
2991 }
2992 std::fs::write(&cfg_path, cfg.to_string()).unwrap();
2993 let s = SessionBuilder::new()
2994 .model(dir.path())
2995 .family("gemma/gemma-4-e2b-it")
2996 .build()
2997 .unwrap();
2998 assert_eq!(s.config().sliding_window, Some(512));
2999 assert_eq!(s.config().partial_rotary_factor, Some(0.25));
3000 assert!(s.config().head_dim.unwrap_or(0) > 0);
3001 assert!(s.config().global_head_dim.unwrap_or(0) > 0);
3002 assert_eq!(
3003 s.config().layer_types.as_ref().map(|t| t.len()),
3004 Some(s.config().num_layers)
3005 );
3006 }
3007
3008 #[test]
3009 fn gemma3_text_not_gemma3n() {
3010 assert!(is_gemma3_text("gemma/gemma-3-270m-it"));
3011 assert!(is_gemma3_text("gemma/gemma-3-1b-it"));
3012 assert!(!is_gemma3_text("gemma/gemma-3n-e2b-it"));
3013 assert!(!is_gemma3_text("gemma/gemma-4-e2b-it"));
3014 let t = default_gemma3_layer_types(18);
3015 assert_eq!(t[5], "full_attention");
3016 assert_eq!(t[11], "full_attention");
3017 assert_eq!(t[17], "full_attention");
3018 assert_eq!(t.iter().filter(|s| *s == "sliding_attention").count(), 15);
3019 }
3020
3021 #[test]
3022 fn session_infers_family_from_hub_cache_dirname() {
3023 let parent = tempfile::tempdir().unwrap();
3024 let bundle = parent.path().join("gemma-3-1b-it_q326");
3025 std::fs::create_dir(&bundle).unwrap();
3026 write_tiny_q4_bundle(&bundle).unwrap();
3027 let s = SessionBuilder::new().model(&bundle).build().unwrap();
3028 assert_eq!(s.family().path(), "gemma/gemma-3-1b-it");
3029
3030 let explicit = SessionBuilder::new()
3031 .model(&bundle)
3032 .family("gemma/gemma-4-e2b-it")
3033 .build()
3034 .unwrap();
3035 assert_eq!(explicit.family().path(), "gemma/gemma-4-e2b-it");
3036 }
3037
3038 #[test]
3039 fn gemma3_fills_hub_bundle_and_dual_rope() {
3040 let dir = tempfile::tempdir().unwrap();
3041 write_tiny_q4_bundle(dir.path()).unwrap();
3042 let cfg_path = dir.path().join("config.json");
3043 let raw = std::fs::read_to_string(&cfg_path).unwrap();
3044 let mut cfg: Value = serde_json::from_str(&raw).unwrap();
3045 let model = cfg["model"].as_object_mut().unwrap();
3046 for key in ["layer_types", "sliding_window", "hidden_act"] {
3047 model.remove(key);
3048 }
3049 std::fs::write(&cfg_path, cfg.to_string()).unwrap();
3050 let mut s = SessionBuilder::new()
3051 .model(dir.path())
3052 .family("gemma/gemma-3-270m-it")
3053 .build()
3054 .unwrap();
3055 assert_eq!(s.config().sliding_window, Some(512));
3056 assert_eq!(s.config().hidden_act.as_deref(), Some("gelu_pytorch_tanh"));
3057 let types = s.config().layer_types.as_ref().expect("layer_types");
3058 assert_eq!(types.len(), s.config().num_layers);
3059 assert!(types.iter().all(|t| t == "sliding_attention"));
3060 assert_eq!(
3061 s.layer_rope_params(AttnKind::Sliding),
3062 (10_000.0, RopeMode::Full)
3063 );
3064 assert_eq!(
3065 s.layer_rope_params(AttnKind::Full),
3066 (1_000_000.0, RopeMode::Full)
3067 );
3068 let gen = s
3069 .generate(
3070 &[1, 2],
3071 &GenerateOpts {
3072 max_tokens: 2,
3073 temperature: 0.0,
3074 },
3075 )
3076 .unwrap();
3077 assert_eq!(gen.tokens.len(), 2);
3078 }
3079
3080 #[test]
3081 fn gemma3n_not_gemma3_text_and_fills_4plus1_dual_rope() {
3082 assert!(is_gemma3n("gemma/gemma-3n-e2b-it"));
3083 assert!(is_gemma3n("gemma/gemma-3n-e4b-it"));
3084 assert!(!is_gemma3n("gemma/gemma-3-270m-it"));
3085 let t = default_gemma4_layer_types(30);
3086 assert_eq!(t[4], "full_attention");
3087 assert_eq!(t[9], "full_attention");
3088 assert_eq!(t[29], "full_attention");
3089 assert_eq!(t.iter().filter(|s| *s == "sliding_attention").count(), 24);
3090 assert_eq!(gemma3n_default_kv_shared(30), 10);
3091 assert_eq!(gemma3n_default_kv_shared(35), 15);
3092 assert_eq!(gemma3n_default_kv_shared(2), 0);
3093 assert!(
3094 (altup_router_input_scale(2048) - 1.0 / 2048.0).abs() < 1e-12,
3095 "HF router_input_scale is 1/hidden, not 1/sqrt(hidden)"
3096 );
3097
3098 let dir = tempfile::tempdir().unwrap();
3099 write_tiny_q4_bundle(dir.path()).unwrap();
3100 let cfg_path = dir.path().join("config.json");
3101 let raw = std::fs::read_to_string(&cfg_path).unwrap();
3102 let mut cfg: Value = serde_json::from_str(&raw).unwrap();
3103 let model = cfg["model"].as_object_mut().unwrap();
3104 for key in ["layer_types", "sliding_window", "hidden_act"] {
3105 model.remove(key);
3106 }
3107 std::fs::write(&cfg_path, cfg.to_string()).unwrap();
3108 let mut s = SessionBuilder::new()
3109 .model(dir.path())
3110 .family("gemma/gemma-3n-e2b-it")
3111 .build()
3112 .unwrap();
3113 assert!(!s.use_gemma_norm, "Gemma-3n RMSNorm is *w, not *(1+w)");
3114 assert!((s.attn_scale(256) - 1.0).abs() < 1e-6);
3115 assert_eq!(s.final_logit_softcap, Some(30.0));
3116 assert_eq!(s.config().sliding_window, Some(512));
3117 assert_eq!(
3118 s.layer_rope_params(AttnKind::Sliding),
3119 (10_000.0, RopeMode::Full)
3120 );
3121 assert_eq!(
3122 s.layer_rope_params(AttnKind::Full),
3123 (1_000_000.0, RopeMode::Full)
3124 );
3125 let gen = s
3126 .generate(
3127 &[1, 2],
3128 &GenerateOpts {
3129 max_tokens: 2,
3130 temperature: 0.0,
3131 },
3132 )
3133 .unwrap();
3134 assert_eq!(gen.tokens.len(), 2);
3135 }
3136
3137 #[test]
3138 fn generate_tokens() {
3139 let dir = tempfile::tempdir().unwrap();
3140 write_tiny_q4_bundle(dir.path()).unwrap();
3141 let mut s = SessionBuilder::new()
3142 .model(dir.path())
3143 .family("gemma/gemma-4-e2b-it")
3144 .build()
3145 .unwrap();
3146 assert_eq!(s.config().sliding_window, Some(512));
3147 assert_eq!(s.config().partial_rotary_factor, Some(0.25));
3148 assert_eq!(s.config().head_dim, Some(16));
3149 assert_eq!(s.config().global_head_dim, Some(16));
3150 assert_eq!(
3151 s.config().layer_types,
3152 Some(vec!["full_attention".into(), "full_attention".into()])
3153 );
3154 assert_eq!(
3155 s.layer_rope_params(AttnKind::Full),
3156 (1_000_000.0, RopeMode::Proportional(0.25))
3157 );
3158 assert_eq!(
3159 s.layer_rope_params(AttnKind::Sliding),
3160 (10_000.0, RopeMode::Full)
3161 );
3162 assert_eq!(s.attn_window(AttnKind::Sliding), Some(512));
3163 assert_eq!(s.attn_window(AttnKind::Full), None);
3164 let prompt = s.encode_text("hi");
3165 let gen = s
3166 .generate(
3167 &prompt,
3168 &GenerateOpts {
3169 max_tokens: 4,
3170 temperature: 0.0,
3171 },
3172 )
3173 .unwrap();
3174 assert!(!gen.tokens.is_empty());
3175 assert!(!gen.text.is_empty());
3176 }
3177
3178 #[test]
3179 fn split_interleaved_q_gate_matches_hf_chunk() {
3180 let mixed = vec![1.0, 2.0, 10.0, 20.0, 3.0, 4.0, 30.0, 40.0];
3182 let (q, g) = split_interleaved_q_gate(&mixed, 1, 2, 2).unwrap();
3183 assert_eq!(q, vec![1.0, 2.0, 3.0, 4.0]);
3184 assert_eq!(g, vec![10.0, 20.0, 30.0, 40.0]);
3185 }
3186
3187 #[test]
3188 fn qwen35_attn_output_gate_generate() {
3189 let dir = tempfile::tempdir().unwrap();
3190 let hidden = 8usize;
3191 let inter = 16usize;
3192 let vocab = 16usize;
3193 let n_heads = 2usize;
3194 let n_kv = 1usize;
3195 let head_dim = 4usize;
3196 let q_dim = n_heads * head_dim;
3197 let k_dim = n_kv * head_dim;
3198
3199 let mut tensors = serde_json::Map::new();
3200 let mut bin = Vec::new();
3201 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
3202 let offset = bin.len();
3203 for &v in data {
3204 bin.extend_from_slice(&v.to_le_bytes());
3205 }
3206 let nbytes = data.len() * 4;
3207 let mut meta = serde_json::Map::new();
3208 meta.insert("kind".into(), json!("raw"));
3209 meta.insert("dtype".into(), json!("f32"));
3210 meta.insert("shape".into(), json!(shape));
3211 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
3212 tensors.insert(name.to_string(), Value::Object(meta));
3213 };
3214 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
3215 add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
3216 let n1 = vec![1.0f32; hidden];
3217 add_raw("model.layers.0.input_layernorm.weight", vec![hidden], &n1);
3218 add_raw(
3219 "model.layers.0.post_attention_layernorm.weight",
3220 vec![hidden],
3221 &n1,
3222 );
3223 let wq = vec![0.02f32; 2 * q_dim * hidden];
3224 let wk = vec![0.02f32; k_dim * hidden];
3225 let wv = vec![0.02f32; k_dim * hidden];
3226 let wo = vec![0.02f32; hidden * q_dim];
3227 add_raw(
3228 "model.layers.0.self_attn.q_proj.weight",
3229 vec![2 * q_dim, hidden],
3230 &wq,
3231 );
3232 add_raw(
3233 "model.layers.0.self_attn.k_proj.weight",
3234 vec![k_dim, hidden],
3235 &wk,
3236 );
3237 add_raw(
3238 "model.layers.0.self_attn.v_proj.weight",
3239 vec![k_dim, hidden],
3240 &wv,
3241 );
3242 add_raw(
3243 "model.layers.0.self_attn.o_proj.weight",
3244 vec![hidden, q_dim],
3245 &wo,
3246 );
3247 let qn = vec![1.0f32; head_dim];
3248 add_raw(
3249 "model.layers.0.self_attn.q_norm.weight",
3250 vec![head_dim],
3251 &qn,
3252 );
3253 add_raw(
3254 "model.layers.0.self_attn.k_norm.weight",
3255 vec![head_dim],
3256 &qn,
3257 );
3258 let g = vec![0.02f32; inter * hidden];
3259 let d = vec![0.02f32; hidden * inter];
3260 add_raw(
3261 "model.layers.0.mlp.gate_proj.weight",
3262 vec![inter, hidden],
3263 &g,
3264 );
3265 add_raw("model.layers.0.mlp.up_proj.weight", vec![inter, hidden], &g);
3266 add_raw(
3267 "model.layers.0.mlp.down_proj.weight",
3268 vec![hidden, inter],
3269 &d,
3270 );
3271 add_raw("model.norm.weight", vec![hidden], &n1);
3272 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
3273 let cfg = json!({
3274 "format": "aria-quant-bundle",
3275 "format_version": 2,
3276 "quantization": "test",
3277 "hadamard_seed": 0,
3278 "model": {
3279 "hidden_size": hidden,
3280 "num_layers": 1,
3281 "num_attention_heads": n_heads,
3282 "num_kv_heads": n_kv,
3283 "head_dim": head_dim,
3284 "intermediate_size": inter,
3285 "vocab_size": vocab,
3286 "context_length": 32,
3287 "rope_theta": 10000.0,
3288 "layer_types": ["full_attention"]
3289 },
3290 "tensors": tensors
3291 });
3292 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
3293 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
3294 let mut s = SessionBuilder::new()
3295 .model(dir.path())
3296 .family("qwen/qwen3-0.6b")
3297 .build()
3298 .unwrap();
3299 let gen = s
3300 .generate(
3301 &[1, 2],
3302 &GenerateOpts {
3303 max_tokens: 2,
3304 temperature: 0.0,
3305 },
3306 )
3307 .unwrap();
3308 assert_eq!(gen.tokens.len(), 2);
3309 }
3310
3311 #[test]
3312 fn materialize_accepts_hf_tensor_names() {
3313 let dir = tempfile::tempdir().unwrap();
3315 let hidden = 8usize;
3316 let layers = 1usize;
3317 let inter = 16usize;
3318 let vocab = 16usize;
3319 let n_heads = 2usize;
3320 let n_kv = 1usize;
3321 let head_dim = 4usize; let q_dim = n_heads * head_dim;
3323 let k_dim = n_kv * head_dim;
3324
3325 let mut tensors = serde_json::Map::new();
3326 let mut bin = Vec::new();
3327 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
3328 let offset = bin.len();
3329 for &v in data {
3330 bin.extend_from_slice(&v.to_le_bytes());
3331 }
3332 let nbytes = data.len() * 4;
3333 let mut meta = serde_json::Map::new();
3334 meta.insert("kind".into(), json!("raw"));
3335 meta.insert("dtype".into(), json!("f32"));
3336 meta.insert("shape".into(), json!(shape));
3337 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
3338 tensors.insert(name.to_string(), Value::Object(meta));
3339 };
3340 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
3341 add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
3342 let n1 = vec![1.0f32; hidden];
3343 add_raw("model.layers.0.input_layernorm.weight", vec![hidden], &n1);
3344 add_raw(
3345 "model.layers.0.post_attention_layernorm.weight",
3346 vec![hidden],
3347 &n1,
3348 );
3349 let wq = vec![0.01f32; q_dim * hidden];
3350 let wk = vec![0.01f32; k_dim * hidden];
3351 let wv = vec![0.01f32; k_dim * hidden];
3352 let wo = vec![0.01f32; hidden * q_dim];
3353 add_raw(
3354 "model.layers.0.self_attn.q_proj.weight",
3355 vec![q_dim, hidden],
3356 &wq,
3357 );
3358 add_raw(
3359 "model.layers.0.self_attn.k_proj.weight",
3360 vec![k_dim, hidden],
3361 &wk,
3362 );
3363 add_raw(
3364 "model.layers.0.self_attn.v_proj.weight",
3365 vec![k_dim, hidden],
3366 &wv,
3367 );
3368 add_raw(
3369 "model.layers.0.self_attn.o_proj.weight",
3370 vec![hidden, q_dim],
3371 &wo,
3372 );
3373 let g = vec![0.01f32; inter * hidden];
3374 let d = vec![0.01f32; hidden * inter];
3375 add_raw(
3376 "model.layers.0.mlp.gate_proj.weight",
3377 vec![inter, hidden],
3378 &g,
3379 );
3380 add_raw("model.layers.0.mlp.up_proj.weight", vec![inter, hidden], &g);
3381 add_raw(
3382 "model.layers.0.mlp.down_proj.weight",
3383 vec![hidden, inter],
3384 &d,
3385 );
3386 add_raw("model.norm.weight", vec![hidden], &n1);
3387 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
3388
3389 let cfg = json!({
3390 "format": "aria-quant-bundle",
3391 "format_version": 2,
3392 "quantization": "test",
3393 "group_size_default": 32,
3394 "hadamard_seed": 0,
3395 "model": {
3396 "hidden_size": hidden,
3397 "num_layers": layers,
3398 "num_attention_heads": n_heads,
3399 "num_kv_heads": n_kv,
3400 "intermediate_size": inter,
3401 "vocab_size": vocab,
3402 "context_length": 32,
3403 "rope_theta": 10000.0
3404 },
3405 "tensors": tensors
3406 });
3407 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
3408 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
3409
3410 let mut s = SessionBuilder::new()
3411 .model(dir.path())
3412 .family("qwen/qwen3-0.6b")
3413 .build()
3414 .unwrap();
3415 let gen = s
3416 .generate(
3417 &[1, 2],
3418 &GenerateOpts {
3419 max_tokens: 2,
3420 temperature: 0.0,
3421 },
3422 )
3423 .unwrap();
3424 assert_eq!(gen.tokens.len(), 2);
3425 }
3426
3427 #[test]
3428 fn materialize_accepts_language_model_prefix_and_pre_ffn_norm() {
3429 let dir = tempfile::tempdir().unwrap();
3431 let hidden = 8usize;
3432 let layers = 1usize;
3433 let inter = 16usize;
3434 let vocab = 16usize;
3435 let n_heads = 2usize;
3436 let n_kv = 1usize;
3437 let head_dim = 4usize;
3438 let q_dim = n_heads * head_dim;
3439 let k_dim = n_kv * head_dim;
3440
3441 let mut tensors = serde_json::Map::new();
3442 let mut bin = Vec::new();
3443 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
3444 let offset = bin.len();
3445 for &v in data {
3446 bin.extend_from_slice(&v.to_le_bytes());
3447 }
3448 let nbytes = data.len() * 4;
3449 let mut meta = serde_json::Map::new();
3450 meta.insert("kind".into(), json!("raw"));
3451 meta.insert("dtype".into(), json!("f32"));
3452 meta.insert("shape".into(), json!(shape));
3453 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
3454 tensors.insert(name.to_string(), Value::Object(meta));
3455 };
3456 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
3457 let p = "model.language_model";
3458 add_raw(
3459 &format!("{p}.embed_tokens.weight"),
3460 vec![vocab, hidden],
3461 &emb,
3462 );
3463 let n1 = vec![1.0f32; hidden];
3464 add_raw(
3465 &format!("{p}.layers.0.input_layernorm.weight"),
3466 vec![hidden],
3467 &n1,
3468 );
3469 add_raw(
3470 &format!("{p}.layers.0.pre_feedforward_layernorm.weight"),
3471 vec![hidden],
3472 &n1,
3473 );
3474 let wq = vec![0.01f32; q_dim * hidden];
3475 let wk = vec![0.01f32; k_dim * hidden];
3476 let wv = vec![0.01f32; k_dim * hidden];
3477 let wo = vec![0.01f32; hidden * q_dim];
3478 add_raw(
3479 &format!("{p}.layers.0.self_attn.q_proj.weight"),
3480 vec![q_dim, hidden],
3481 &wq,
3482 );
3483 add_raw(
3484 &format!("{p}.layers.0.self_attn.k_proj.weight"),
3485 vec![k_dim, hidden],
3486 &wk,
3487 );
3488 add_raw(
3489 &format!("{p}.layers.0.self_attn.v_proj.weight"),
3490 vec![k_dim, hidden],
3491 &wv,
3492 );
3493 add_raw(
3494 &format!("{p}.layers.0.self_attn.o_proj.weight"),
3495 vec![hidden, q_dim],
3496 &wo,
3497 );
3498 let g = vec![0.01f32; inter * hidden];
3499 let d = vec![0.01f32; hidden * inter];
3500 add_raw(
3501 &format!("{p}.layers.0.mlp.gate_proj.weight"),
3502 vec![inter, hidden],
3503 &g,
3504 );
3505 add_raw(
3506 &format!("{p}.layers.0.mlp.up_proj.weight"),
3507 vec![inter, hidden],
3508 &g,
3509 );
3510 add_raw(
3511 &format!("{p}.layers.0.mlp.down_proj.weight"),
3512 vec![hidden, inter],
3513 &d,
3514 );
3515 add_raw(&format!("{p}.norm.weight"), vec![hidden], &n1);
3516 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
3517
3518 let cfg = json!({
3519 "format": "aria-quant-bundle",
3520 "format_version": 2,
3521 "quantization": "test",
3522 "group_size_default": 32,
3523 "hadamard_seed": 0,
3524 "model": {
3525 "hidden_size": hidden,
3526 "num_layers": layers,
3527 "num_attention_heads": n_heads,
3528 "num_kv_heads": n_kv,
3529 "intermediate_size": inter,
3530 "vocab_size": vocab,
3531 "context_length": 32,
3532 "rope_theta": 10000.0,
3533 "head_dim": head_dim,
3534 "global_head_dim": head_dim,
3535 "sliding_window": 512,
3536 "partial_rotary_factor": 0.25,
3537 "layer_types": ["full_attention"]
3538 },
3539 "tensors": tensors
3540 });
3541 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
3542 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
3543
3544 let mut s = SessionBuilder::new()
3545 .model(dir.path())
3546 .family("gemma/gemma-4-e2b-it")
3547 .build()
3548 .unwrap();
3549 let gen = s
3550 .generate(
3551 &[1, 2],
3552 &GenerateOpts {
3553 max_tokens: 2,
3554 temperature: 0.0,
3555 },
3556 )
3557 .unwrap();
3558 assert_eq!(gen.tokens.len(), 2);
3559 }
3560
3561 #[test]
3562 fn gemma4_style_double_wide_mlp_and_shared_kv() {
3563 let dir = tempfile::tempdir().unwrap();
3566 let hidden = 8usize;
3567 let layers = 2usize;
3568 let inter = 16usize;
3569 let inter_wide = 32usize;
3570 let vocab = 16usize;
3571 let n_heads = 2usize;
3572 let n_kv = 1usize;
3573 let head_dim = 4usize;
3574 let q_dim = n_heads * head_dim;
3575 let k_dim = n_kv * head_dim;
3576 let p = "model.language_model";
3577
3578 let mut tensors = serde_json::Map::new();
3579 let mut bin = Vec::new();
3580 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
3581 let offset = bin.len();
3582 for &v in data {
3583 bin.extend_from_slice(&v.to_le_bytes());
3584 }
3585 let nbytes = data.len() * 4;
3586 let mut meta = serde_json::Map::new();
3587 meta.insert("kind".into(), json!("raw"));
3588 meta.insert("dtype".into(), json!("f32"));
3589 meta.insert("shape".into(), json!(shape));
3590 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
3591 tensors.insert(name.to_string(), Value::Object(meta));
3592 };
3593 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
3594 add_raw(
3595 &format!("{p}.embed_tokens.weight"),
3596 vec![vocab, hidden],
3597 &emb,
3598 );
3599 let n1 = vec![1.0f32; hidden];
3600 let wq = vec![0.01f32; q_dim * hidden];
3601 let wk = vec![0.01f32; k_dim * hidden];
3602 let wv = vec![0.01f32; k_dim * hidden];
3603 let wo = vec![0.01f32; hidden * q_dim];
3604 for li in 0..layers {
3605 let layer_inter = if li == 0 { inter } else { inter_wide };
3606 add_raw(
3607 &format!("{p}.layers.{li}.input_layernorm.weight"),
3608 vec![hidden],
3609 &n1,
3610 );
3611 add_raw(
3612 &format!("{p}.layers.{li}.pre_feedforward_layernorm.weight"),
3613 vec![hidden],
3614 &n1,
3615 );
3616 add_raw(
3617 &format!("{p}.layers.{li}.self_attn.q_proj.weight"),
3618 vec![q_dim, hidden],
3619 &wq,
3620 );
3621 if li == 0 {
3622 add_raw(
3623 &format!("{p}.layers.{li}.self_attn.k_proj.weight"),
3624 vec![k_dim, hidden],
3625 &wk,
3626 );
3627 add_raw(
3628 &format!("{p}.layers.{li}.self_attn.v_proj.weight"),
3629 vec![k_dim, hidden],
3630 &wv,
3631 );
3632 }
3633 add_raw(
3634 &format!("{p}.layers.{li}.self_attn.o_proj.weight"),
3635 vec![hidden, q_dim],
3636 &wo,
3637 );
3638 let g = vec![0.01f32; layer_inter * hidden];
3639 let d = vec![0.01f32; hidden * layer_inter];
3640 add_raw(
3641 &format!("{p}.layers.{li}.mlp.gate_proj.weight"),
3642 vec![layer_inter, hidden],
3643 &g,
3644 );
3645 add_raw(
3646 &format!("{p}.layers.{li}.mlp.up_proj.weight"),
3647 vec![layer_inter, hidden],
3648 &g,
3649 );
3650 add_raw(
3651 &format!("{p}.layers.{li}.mlp.down_proj.weight"),
3652 vec![hidden, layer_inter],
3653 &d,
3654 );
3655 }
3656 add_raw(&format!("{p}.norm.weight"), vec![hidden], &n1);
3657 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
3658
3659 let cfg = json!({
3660 "format": "aria-quant-bundle",
3661 "format_version": 2,
3662 "quantization": "test",
3663 "group_size_default": 32,
3664 "hadamard_seed": 0,
3665 "model": {
3666 "hidden_size": hidden,
3667 "num_layers": layers,
3668 "num_attention_heads": n_heads,
3669 "num_kv_heads": n_kv,
3670 "intermediate_size": inter,
3671 "vocab_size": vocab,
3672 "context_length": 32,
3673 "rope_theta": 10000.0,
3674 "num_kv_shared_layers": 1,
3675 "head_dim": head_dim,
3676 "global_head_dim": head_dim,
3677 "sliding_window": 512,
3678 "partial_rotary_factor": 0.25,
3679 "layer_types": ["full_attention", "full_attention"]
3680 },
3681 "tensors": tensors
3682 });
3683 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
3684 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
3685
3686 let mut s = SessionBuilder::new()
3687 .model(dir.path())
3688 .family("gemma/gemma-4-e2b-it")
3689 .build()
3690 .unwrap();
3691 let gen = s
3692 .generate(
3693 &[1, 2],
3694 &GenerateOpts {
3695 max_tokens: 2,
3696 temperature: 0.0,
3697 },
3698 )
3699 .unwrap();
3700 assert_eq!(gen.tokens.len(), 2);
3701 }
3702
3703 #[test]
3704 fn stage_b_arch_classes_generate() {
3705 for (path, arch) in arch_class_representatives() {
3706 if matches!(arch, ArchClass::VL | ArchClass::VLA | ArchClass::TextMoE) {
3707 continue; }
3709 if path.contains("qwen3.5") || path.contains("bonsai") {
3711 let dir = tempfile::tempdir().unwrap();
3712 write_tiny_q4_bundle(dir.path()).unwrap();
3713 let err = SessionBuilder::new()
3714 .model(dir.path())
3715 .family(*path)
3716 .build()
3717 .unwrap_err();
3718 assert!(
3719 matches!(err, EngineError::Unsupported(_)),
3720 "{path}: {err:?}"
3721 );
3722 continue;
3723 }
3724 assert!(require_stage_b(path).is_ok(), "{path}");
3725 let dir = tempfile::tempdir().unwrap();
3726 write_tiny_q4_bundle(dir.path()).unwrap();
3727 let mut s = SessionBuilder::new()
3728 .model(dir.path())
3729 .family(*path)
3730 .build()
3731 .unwrap();
3732 assert_eq!(s.arch(), *arch);
3733 assert!(!s.graph_hook_name().is_empty());
3734 let gen = s
3735 .generate(
3736 &s.encode_text("ok"),
3737 &GenerateOpts {
3738 max_tokens: 2,
3739 temperature: 0.0,
3740 },
3741 )
3742 .unwrap();
3743 assert!(!gen.tokens.is_empty(), "{path}");
3744 }
3745 }
3746
3747 #[test]
3748 fn stage_c_vl_vla_hooks() {
3749 let dir = tempfile::tempdir().unwrap();
3750 write_tiny_q4_bundle(dir.path()).unwrap();
3751 let s = SessionBuilder::new()
3752 .model(dir.path())
3753 .family("lfm/lfm2-vl-450m")
3754 .build()
3755 .unwrap();
3756 let rgb = vec![10u8; 3 * 4 * 4];
3757 let err = s.vision_prefix(&rgb, 4, 4).unwrap_err();
3758 assert!(matches!(err, EngineError::Unsupported(_)));
3759
3760 let vla = SessionBuilder::new()
3761 .model(dir.path())
3762 .family("openvla/openvla-7b")
3763 .build()
3764 .unwrap();
3765 let err = vla.predict_action("move", 7).unwrap_err();
3766 assert!(matches!(err, EngineError::Unsupported(_)));
3767 let emb = vla.embed_text("hello").unwrap();
3768 assert_eq!(emb.len(), vla.config().hidden_size);
3769 }
3770
3771 #[test]
3772 fn unknown_family() {
3773 let err = SessionBuilder::new()
3774 .model("/tmp")
3775 .family("no/such-model")
3776 .build()
3777 .unwrap_err();
3778 assert!(matches!(err, EngineError::UnsupportedFamily(_)));
3779 }
3780
3781 #[test]
3782 fn greedy_deterministic() {
3783 let dir = tempfile::tempdir().unwrap();
3784 write_tiny_q4_bundle(dir.path()).unwrap();
3785 let mut s = SessionBuilder::new()
3786 .model(dir.path())
3787 .family("gemma/gemma-4-e2b-it")
3788 .build()
3789 .unwrap();
3790 let prompt = s.encode_text("hi");
3791 let opts = GenerateOpts {
3792 max_tokens: 3,
3793 temperature: 0.0,
3794 };
3795 let a = s.generate(&prompt, &opts).unwrap();
3796 let b = s.generate(&prompt, &opts).unwrap();
3797 assert_eq!(a.tokens, b.tokens);
3798 assert_eq!(a.tokens.len(), 3);
3799 }
3800
3801 #[test]
3802 fn encode_chat_is_longer_than_raw_user_text() {
3803 let dir = tempfile::tempdir().unwrap();
3804 write_tiny_q4_bundle(dir.path()).unwrap();
3805 let s = SessionBuilder::new()
3806 .model(dir.path())
3807 .family("qwen/qwen3-0.6b")
3808 .build()
3809 .unwrap();
3810 let raw = s.encode_text("Hello");
3811 let chat = s.encode_chat(&[ChatTurn::new("user", "Hello")]);
3812 assert!(
3813 chat.len() > raw.len(),
3814 "chat template should wrap the user turn (raw={}, chat={})",
3815 raw.len(),
3816 chat.len()
3817 );
3818 assert!(
3819 (s.config().rope_theta - 1_000_000.0).abs() < 1.0,
3820 "Qwen3 must not keep Llama-default rope_theta=10000, got {}",
3821 s.config().rope_theta
3822 );
3823 }
3824
3825 #[test]
3826 fn incremental_decode_matches_full_recompute() {
3827 let dir = tempfile::tempdir().unwrap();
3828 write_tiny_q4_bundle(dir.path()).unwrap();
3829 let mut s = SessionBuilder::new()
3830 .model(dir.path())
3831 .family("gemma/gemma-4-e2b-it")
3832 .build()
3833 .unwrap();
3834 let prompt = s.encode_text("hi");
3835 let max_tokens = 5usize;
3836
3837 let mut prefix = prompt.clone();
3839 if prefix.is_empty() {
3840 prefix.push(1);
3841 }
3842 let mut full_tokens = Vec::new();
3843 for _ in 0..max_tokens {
3844 let logits = s.forward(&prefix).unwrap();
3845 let next = argmax(&logits);
3846 full_tokens.push(next);
3847 prefix.push(next);
3848 if s.is_stop_id(next) {
3849 full_tokens.pop();
3850 break;
3851 }
3852 }
3853
3854 let incr = s
3855 .generate(
3856 &prompt,
3857 &GenerateOpts {
3858 max_tokens,
3859 temperature: 0.0,
3860 },
3861 )
3862 .unwrap();
3863 assert_eq!(
3864 incr.tokens, full_tokens,
3865 "incremental decode must match full-recompute greedy tokens"
3866 );
3867 }
3868
3869 #[test]
3870 fn profile_records_load_and_generate() {
3871 let dir = tempfile::tempdir().unwrap();
3872 write_tiny_q4_bundle(dir.path()).unwrap();
3873 let mut s = SessionBuilder::new()
3874 .model(dir.path())
3875 .family("gemma/gemma-4-e2b-it")
3876 .compute(ComputePref::Cpu)
3877 .profile(true)
3878 .build()
3879 .unwrap();
3880 assert!(s.compute_label().contains("cpu"));
3881 let load = s.last_profile().expect("load profile");
3882 assert!(!load.ci_fail);
3883 assert!(load.load.materialize_ms >= 0.0);
3884 s.generate(
3885 &s.encode_text("hi"),
3886 &GenerateOpts {
3887 max_tokens: 2,
3888 temperature: 0.0,
3889 },
3890 )
3891 .unwrap();
3892 let p = s.last_profile().expect("generate profile");
3893 let g = p.generate.as_ref().expect("generate timings");
3894 assert!(g.prefill_ms >= 0.0);
3895 assert!(g.decode_ms >= 0.0);
3896 }
3897
3898 #[test]
3899 fn cuda_greedy_matches_cpu_if_available() {
3900 if resolve_compute(ComputePref::Cuda).is_err() {
3901 return;
3902 }
3903 let dir = tempfile::tempdir().unwrap();
3904 write_tiny_q4_bundle(dir.path()).unwrap();
3905 let prompt_text = "hi";
3906 let opts = GenerateOpts {
3907 max_tokens: 4,
3908 temperature: 0.0,
3909 };
3910 let mut cpu = SessionBuilder::new()
3911 .model(dir.path())
3912 .family("gemma/gemma-4-e2b-it")
3913 .compute(ComputePref::Cpu)
3914 .build()
3915 .unwrap();
3916 let mut gpu = SessionBuilder::new()
3917 .model(dir.path())
3918 .family("gemma/gemma-4-e2b-it")
3919 .compute(ComputePref::Cuda)
3920 .build()
3921 .unwrap();
3922 assert!(gpu.compute_label().contains("cuda"));
3923 let prompt = cpu.encode_text(prompt_text);
3924 let a = cpu.generate(&prompt, &opts).unwrap();
3925 let b = gpu.generate(&prompt, &opts).unwrap();
3926 assert_eq!(
3927 a.tokens, b.tokens,
3928 "CUDA greedy tokens must match CPU (tiny bundle)"
3929 );
3930 }
3931
3932 #[test]
3933 fn max_tokens_zero_rejected() {
3934 let dir = tempfile::tempdir().unwrap();
3935 write_tiny_q4_bundle(dir.path()).unwrap();
3936 let mut s = SessionBuilder::new()
3937 .model(dir.path())
3938 .family("gemma/gemma-4-e2b-it")
3939 .build()
3940 .unwrap();
3941 let err = s
3942 .generate(
3943 &s.encode_text("x"),
3944 &GenerateOpts {
3945 max_tokens: 0,
3946 temperature: 0.0,
3947 },
3948 )
3949 .unwrap_err();
3950 assert!(matches!(err, EngineError::InvalidParam(_)));
3951 }
3952
3953 #[test]
3954 fn moe_family_refuses_dense_stub() {
3955 let dir = tempfile::tempdir().unwrap();
3956 write_tiny_q4_bundle(dir.path()).unwrap();
3957 let err = SessionBuilder::new()
3958 .model(dir.path())
3959 .family("lfm/lfm2-8b-a1b")
3960 .build()
3961 .unwrap_err();
3962 assert!(matches!(err, EngineError::Unsupported(_)));
3963 assert_eq!(
3964 lookup_family("lfm/lfm2-8b-a1b").unwrap().arch,
3965 ArchClass::TextMoE
3966 );
3967 assert_eq!(graph_hook(ArchClass::TextMoE), "text_moe_decoder");
3968 }
3969
3970 #[test]
3971 fn geometry_gates_conv_and_experts() {
3972 let dir = tempfile::tempdir().unwrap();
3974 write_tiny_q4_bundle(dir.path()).unwrap();
3975 let cfg_path = dir.path().join("config.json");
3976 let mut cfg: Value =
3977 serde_json::from_str(&std::fs::read_to_string(&cfg_path).unwrap()).unwrap();
3978 cfg["model"]["layer_types"] = json!(["conv", "full_attention"]);
3979 std::fs::write(&cfg_path, cfg.to_string()).unwrap();
3980 let err = SessionBuilder::new()
3981 .model(dir.path())
3982 .family("lfm/lfm2-350m")
3983 .build()
3984 .unwrap_err();
3985 assert!(
3986 matches!(err, EngineError::Format(_)),
3987 "expected missing conv tensors, got {err:?}"
3988 );
3989
3990 let dir2 = tempfile::tempdir().unwrap();
3992 write_tiny_q4_bundle(dir2.path()).unwrap();
3993 let cfg_path = dir2.path().join("config.json");
3994 let mut cfg: Value =
3995 serde_json::from_str(&std::fs::read_to_string(&cfg_path).unwrap()).unwrap();
3996 cfg["model"]["layer_types"] = json!(["linear_attention", "full_attention"]);
3997 std::fs::write(&cfg_path, cfg.to_string()).unwrap();
3998 let err = SessionBuilder::new()
3999 .model(dir2.path())
4000 .family("gemma/gemma-3-270m-it")
4001 .build()
4002 .unwrap_err();
4003 assert!(
4004 matches!(err, EngineError::Format(_)),
4005 "expected missing DeltaNet tensors, got {err:?}"
4006 );
4007 }
4008
4009 #[test]
4010 fn lfm_short_conv_and_attn_generate() {
4011 let dir = tempfile::tempdir().unwrap();
4012 let hidden = 8usize;
4013 let layers = 2usize;
4014 let inter = 16usize;
4015 let vocab = 16usize;
4016 let n_heads = 2usize;
4017 let n_kv = 1usize;
4018 let head_dim = 4usize;
4019 let q_dim = n_heads * head_dim;
4020 let k_dim = n_kv * head_dim;
4021 let kernel = 3usize;
4022
4023 let mut tensors = serde_json::Map::new();
4024 let mut bin = Vec::new();
4025 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
4026 let offset = bin.len();
4027 for &v in data {
4028 bin.extend_from_slice(&v.to_le_bytes());
4029 }
4030 let nbytes = data.len() * 4;
4031 let mut meta = serde_json::Map::new();
4032 meta.insert("kind".into(), json!("raw"));
4033 meta.insert("dtype".into(), json!("f32"));
4034 meta.insert("shape".into(), json!(shape));
4035 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
4036 tensors.insert(name.to_string(), Value::Object(meta));
4037 };
4038 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
4039 add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
4040 let n1 = vec![1.0f32; hidden];
4041 add_raw("model.layers.0.operator_norm.weight", vec![hidden], &n1);
4043 add_raw("model.layers.0.ffn_norm.weight", vec![hidden], &n1);
4044 let in_proj = vec![0.02f32; 3 * hidden * hidden];
4045 let out_proj = vec![0.02f32; hidden * hidden];
4046 let conv_w = vec![0.1f32; hidden * kernel];
4047 add_raw(
4048 "model.layers.0.conv.in_proj.weight",
4049 vec![3 * hidden, hidden],
4050 &in_proj,
4051 );
4052 add_raw(
4053 "model.layers.0.conv.out_proj.weight",
4054 vec![hidden, hidden],
4055 &out_proj,
4056 );
4057 add_raw(
4058 "model.layers.0.conv.conv.weight",
4059 vec![hidden, kernel],
4060 &conv_w,
4061 );
4062 let g = vec![0.02f32; inter * hidden];
4063 let d = vec![0.02f32; hidden * inter];
4064 add_raw(
4065 "model.layers.0.mlp.gate_proj.weight",
4066 vec![inter, hidden],
4067 &g,
4068 );
4069 add_raw("model.layers.0.mlp.up_proj.weight", vec![inter, hidden], &g);
4070 add_raw(
4071 "model.layers.0.mlp.down_proj.weight",
4072 vec![hidden, inter],
4073 &d,
4074 );
4075 add_raw("model.layers.1.operator_norm.weight", vec![hidden], &n1);
4077 add_raw(
4078 "model.layers.1.post_attention_layernorm.weight",
4079 vec![hidden],
4080 &n1,
4081 );
4082 let wq = vec![0.02f32; q_dim * hidden];
4083 let wk = vec![0.02f32; k_dim * hidden];
4084 let wv = vec![0.02f32; k_dim * hidden];
4085 let wo = vec![0.02f32; hidden * q_dim];
4086 add_raw(
4087 "model.layers.1.self_attn.q_proj.weight",
4088 vec![q_dim, hidden],
4089 &wq,
4090 );
4091 add_raw(
4092 "model.layers.1.self_attn.k_proj.weight",
4093 vec![k_dim, hidden],
4094 &wk,
4095 );
4096 add_raw(
4097 "model.layers.1.self_attn.v_proj.weight",
4098 vec![k_dim, hidden],
4099 &wv,
4100 );
4101 add_raw(
4102 "model.layers.1.self_attn.o_proj.weight",
4103 vec![hidden, q_dim],
4104 &wo,
4105 );
4106 add_raw(
4107 "model.layers.1.mlp.gate_proj.weight",
4108 vec![inter, hidden],
4109 &g,
4110 );
4111 add_raw("model.layers.1.mlp.up_proj.weight", vec![inter, hidden], &g);
4112 add_raw(
4113 "model.layers.1.mlp.down_proj.weight",
4114 vec![hidden, inter],
4115 &d,
4116 );
4117 add_raw("model.norm.weight", vec![hidden], &n1);
4118 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
4119
4120 let cfg = json!({
4121 "format": "aria-quant-bundle",
4122 "format_version": 2,
4123 "quantization": "test",
4124 "group_size_default": 32,
4125 "hadamard_seed": 0,
4126 "model": {
4127 "hidden_size": hidden,
4128 "num_layers": layers,
4129 "num_attention_heads": n_heads,
4130 "num_kv_heads": n_kv,
4131 "head_dim": head_dim,
4132 "intermediate_size": inter,
4133 "vocab_size": vocab,
4134 "context_length": 32,
4135 "rope_theta": 10000.0,
4136 "conv_l_cache": kernel,
4137 "layer_types": ["conv", "full_attention"]
4138 },
4139 "tensors": tensors
4140 });
4141 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
4142 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
4143
4144 let mut s = SessionBuilder::new()
4145 .model(dir.path())
4146 .family("lfm/lfm2-350m")
4147 .build()
4148 .unwrap();
4149 let gen = s
4150 .generate(
4151 &[1, 2, 3],
4152 &GenerateOpts {
4153 max_tokens: 2,
4154 temperature: 0.0,
4155 },
4156 )
4157 .unwrap();
4158 assert_eq!(gen.tokens.len(), 2);
4159 }
4160
4161 #[test]
4162 fn moe_topk_experts_generate() {
4163 let dir = tempfile::tempdir().unwrap();
4164 let hidden = 8usize;
4165 let layers = 1usize;
4166 let inter = 16usize;
4167 let vocab = 16usize;
4168 let n_heads = 2usize;
4169 let n_kv = 1usize;
4170 let head_dim = 4usize;
4171 let q_dim = n_heads * head_dim;
4172 let k_dim = n_kv * head_dim;
4173 let n_experts = 4usize;
4174 let top_k = 2usize;
4175
4176 let mut tensors = serde_json::Map::new();
4177 let mut bin = Vec::new();
4178 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
4179 let offset = bin.len();
4180 for &v in data {
4181 bin.extend_from_slice(&v.to_le_bytes());
4182 }
4183 let nbytes = data.len() * 4;
4184 let mut meta = serde_json::Map::new();
4185 meta.insert("kind".into(), json!("raw"));
4186 meta.insert("dtype".into(), json!("f32"));
4187 meta.insert("shape".into(), json!(shape));
4188 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
4189 tensors.insert(name.to_string(), Value::Object(meta));
4190 };
4191 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
4192 add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
4193 let n1 = vec![1.0f32; hidden];
4194 add_raw("model.layers.0.input_layernorm.weight", vec![hidden], &n1);
4195 add_raw(
4196 "model.layers.0.post_attention_layernorm.weight",
4197 vec![hidden],
4198 &n1,
4199 );
4200 let wq = vec![0.02f32; q_dim * hidden];
4201 let wk = vec![0.02f32; k_dim * hidden];
4202 let wv = vec![0.02f32; k_dim * hidden];
4203 let wo = vec![0.02f32; hidden * q_dim];
4204 add_raw(
4205 "model.layers.0.self_attn.q_proj.weight",
4206 vec![q_dim, hidden],
4207 &wq,
4208 );
4209 add_raw(
4210 "model.layers.0.self_attn.k_proj.weight",
4211 vec![k_dim, hidden],
4212 &wk,
4213 );
4214 add_raw(
4215 "model.layers.0.self_attn.v_proj.weight",
4216 vec![k_dim, hidden],
4217 &wv,
4218 );
4219 add_raw(
4220 "model.layers.0.self_attn.o_proj.weight",
4221 vec![hidden, q_dim],
4222 &wo,
4223 );
4224 let router: Vec<f32> = (0..n_experts * hidden)
4225 .map(|i| ((i % n_experts) as f32) * 0.1)
4226 .collect();
4227 add_raw(
4228 "model.layers.0.block_sparse_moe.gate.weight",
4229 vec![n_experts, hidden],
4230 &router,
4231 );
4232 let g = vec![0.02f32; inter * hidden];
4233 let d = vec![0.02f32; hidden * inter];
4234 for e in 0..n_experts {
4235 add_raw(
4236 &format!("model.layers.0.block_sparse_moe.experts.{e}.w1.weight"),
4237 vec![inter, hidden],
4238 &g,
4239 );
4240 add_raw(
4241 &format!("model.layers.0.block_sparse_moe.experts.{e}.w3.weight"),
4242 vec![inter, hidden],
4243 &g,
4244 );
4245 add_raw(
4246 &format!("model.layers.0.block_sparse_moe.experts.{e}.w2.weight"),
4247 vec![hidden, inter],
4248 &d,
4249 );
4250 }
4251 add_raw("model.norm.weight", vec![hidden], &n1);
4252 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
4253
4254 let cfg = json!({
4255 "format": "aria-quant-bundle",
4256 "format_version": 2,
4257 "quantization": "test",
4258 "group_size_default": 32,
4259 "hadamard_seed": 0,
4260 "model": {
4261 "hidden_size": hidden,
4262 "num_layers": layers,
4263 "num_attention_heads": n_heads,
4264 "num_kv_heads": n_kv,
4265 "head_dim": head_dim,
4266 "intermediate_size": inter,
4267 "vocab_size": vocab,
4268 "context_length": 32,
4269 "rope_theta": 10000.0,
4270 "num_experts": n_experts,
4271 "num_experts_per_tok": top_k
4272 },
4273 "tensors": tensors
4274 });
4275 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
4276 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
4277
4278 let mut s = SessionBuilder::new()
4279 .model(dir.path())
4280 .family("inkling/inkling-small")
4281 .build()
4282 .unwrap();
4283 assert_eq!(s.arch(), ArchClass::TextMoE);
4284 assert_eq!(s.graph_hook_name(), "text_moe_decoder");
4285 let gen = s
4286 .generate(
4287 &[1, 2],
4288 &GenerateOpts {
4289 max_tokens: 2,
4290 temperature: 0.0,
4291 },
4292 )
4293 .unwrap();
4294 assert_eq!(gen.tokens.len(), 2);
4295 }
4296
4297 #[test]
4298 fn tiny_q4_codebook_weights_unrotate_on_load() {
4299 let dir = tempfile::tempdir().unwrap();
4300 write_tiny_q4_bundle(dir.path()).unwrap();
4301 let b = load_bundle(dir.path()).unwrap();
4302 let w = b.weight_loaded("blk.0.attn_q.weight").unwrap();
4303 assert!(
4304 w.hdm_seed.is_none(),
4305 "reconstruct_weight path stores original-space W for linear()"
4306 );
4307 let mut s = SessionBuilder::new()
4308 .model(dir.path())
4309 .family("gemma/gemma-4-e2b-it")
4310 .build()
4311 .unwrap();
4312 let gen = s
4313 .generate(
4314 &[1, 2],
4315 &GenerateOpts {
4316 max_tokens: 2,
4317 temperature: 0.0,
4318 },
4319 )
4320 .unwrap();
4321 assert_eq!(gen.tokens.len(), 2);
4322 }
4323
4324 #[test]
4325 fn gemma_hidden_act_geglu_and_qk_norm() {
4326 let dir = tempfile::tempdir().unwrap();
4327 let hidden = 8usize;
4328 let layers = 1usize;
4329 let inter = 16usize;
4330 let vocab = 16usize;
4331 let n_heads = 2usize;
4332 let n_kv = 1usize;
4333 let head_dim = 4usize;
4334 let q_dim = n_heads * head_dim;
4335 let k_dim = n_kv * head_dim;
4336 let p = "model.language_model";
4337
4338 let mut tensors = serde_json::Map::new();
4339 let mut bin = Vec::new();
4340 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
4341 let offset = bin.len();
4342 for &v in data {
4343 bin.extend_from_slice(&v.to_le_bytes());
4344 }
4345 let nbytes = data.len() * 4;
4346 let mut meta = serde_json::Map::new();
4347 meta.insert("kind".into(), json!("raw"));
4348 meta.insert("dtype".into(), json!("f32"));
4349 meta.insert("shape".into(), json!(shape));
4350 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
4351 tensors.insert(name.to_string(), Value::Object(meta));
4352 };
4353 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
4354 add_raw(
4355 &format!("{p}.embed_tokens.weight"),
4356 vec![vocab, hidden],
4357 &emb,
4358 );
4359 let n1 = vec![1.0f32; hidden];
4360 let qn = vec![1.0f32; head_dim];
4361 let kn = vec![1.0f32; head_dim];
4362 let wq = vec![0.01f32; q_dim * hidden];
4363 let wk = vec![0.01f32; k_dim * hidden];
4364 let wv = vec![0.01f32; k_dim * hidden];
4365 let wo = vec![0.01f32; hidden * q_dim];
4366 add_raw(
4367 &format!("{p}.layers.0.input_layernorm.weight"),
4368 vec![hidden],
4369 &n1,
4370 );
4371 add_raw(
4372 &format!("{p}.layers.0.pre_feedforward_layernorm.weight"),
4373 vec![hidden],
4374 &n1,
4375 );
4376 add_raw(
4377 &format!("{p}.layers.0.self_attn.q_proj.weight"),
4378 vec![q_dim, hidden],
4379 &wq,
4380 );
4381 add_raw(
4382 &format!("{p}.layers.0.self_attn.k_proj.weight"),
4383 vec![k_dim, hidden],
4384 &wk,
4385 );
4386 add_raw(
4387 &format!("{p}.layers.0.self_attn.v_proj.weight"),
4388 vec![k_dim, hidden],
4389 &wv,
4390 );
4391 add_raw(
4392 &format!("{p}.layers.0.self_attn.o_proj.weight"),
4393 vec![hidden, q_dim],
4394 &wo,
4395 );
4396 add_raw(
4397 &format!("{p}.layers.0.self_attn.q_norm.weight"),
4398 vec![head_dim],
4399 &qn,
4400 );
4401 add_raw(
4402 &format!("{p}.layers.0.self_attn.k_norm.weight"),
4403 vec![head_dim],
4404 &kn,
4405 );
4406 let g = vec![0.01f32; inter * hidden];
4407 let d = vec![0.01f32; hidden * inter];
4408 add_raw(
4409 &format!("{p}.layers.0.mlp.gate_proj.weight"),
4410 vec![inter, hidden],
4411 &g,
4412 );
4413 add_raw(
4414 &format!("{p}.layers.0.mlp.up_proj.weight"),
4415 vec![inter, hidden],
4416 &g,
4417 );
4418 add_raw(
4419 &format!("{p}.layers.0.mlp.down_proj.weight"),
4420 vec![hidden, inter],
4421 &d,
4422 );
4423 add_raw(&format!("{p}.norm.weight"), vec![hidden], &n1);
4424 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
4425
4426 let cfg = json!({
4427 "format": "aria-quant-bundle",
4428 "format_version": 2,
4429 "quantization": "test",
4430 "group_size_default": 32,
4431 "hadamard_seed": 0,
4432 "model": {
4433 "hidden_size": hidden,
4434 "num_layers": layers,
4435 "num_attention_heads": n_heads,
4436 "num_kv_heads": n_kv,
4437 "head_dim": head_dim,
4438 "global_head_dim": head_dim,
4439 "sliding_window": 512,
4440 "partial_rotary_factor": 0.25,
4441 "intermediate_size": inter,
4442 "vocab_size": vocab,
4443 "context_length": 32,
4444 "rope_theta": 10000.0,
4445 "hidden_act": "gelu_pytorch_tanh",
4446 "layer_types": ["full_attention"]
4447 },
4448 "tensors": tensors
4449 });
4450 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
4451 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
4452
4453 let mut s = SessionBuilder::new()
4454 .model(dir.path())
4455 .family("gemma/gemma-4-e2b-it")
4456 .build()
4457 .unwrap();
4458 assert_eq!(s.config().hidden_act.as_deref(), Some("gelu_pytorch_tanh"));
4459 let gen = s
4460 .generate(
4461 &[1, 2],
4462 &GenerateOpts {
4463 max_tokens: 2,
4464 temperature: 0.0,
4465 },
4466 )
4467 .unwrap();
4468 assert_eq!(gen.tokens.len(), 2);
4469 }
4470
4471 #[test]
4472 fn gated_deltanet_and_full_attn_generate() {
4473 let dir = tempfile::tempdir().unwrap();
4474 let hidden = 8usize;
4475 let inter = 16usize;
4476 let vocab = 16usize;
4477 let n_heads = 2usize;
4478 let n_kv = 1usize;
4479 let head_dim = 4usize;
4480 let q_dim = n_heads * head_dim;
4481 let k_dim = n_kv * head_dim;
4482 let n_lin = 2usize;
4483 let hk = 4usize;
4484 let hv = 4usize;
4485 let key_dim = n_lin * hk;
4486 let value_dim = n_lin * hv;
4487 let conv_k = 4usize;
4488 let conv_dim = key_dim * 2 + value_dim;
4489
4490 let mut tensors = serde_json::Map::new();
4491 let mut bin = Vec::new();
4492 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
4493 let offset = bin.len();
4494 for &v in data {
4495 bin.extend_from_slice(&v.to_le_bytes());
4496 }
4497 let nbytes = data.len() * 4;
4498 let mut meta = serde_json::Map::new();
4499 meta.insert("kind".into(), json!("raw"));
4500 meta.insert("dtype".into(), json!("f32"));
4501 meta.insert("shape".into(), json!(shape));
4502 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
4503 tensors.insert(name.to_string(), Value::Object(meta));
4504 };
4505 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
4506 add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
4507 let n1 = vec![1.0f32; hidden];
4508 add_raw("model.layers.0.input_layernorm.weight", vec![hidden], &n1);
4509 add_raw(
4510 "model.layers.0.post_attention_layernorm.weight",
4511 vec![hidden],
4512 &n1,
4513 );
4514 let qkvz = vec![0.02f32; (2 * key_dim + 2 * value_dim) * hidden];
4515 let ba = vec![0.1f32; 2 * n_lin * hidden];
4516 let conv = vec![0.05f32; conv_dim * conv_k];
4517 let a_log = vec![0.5f32; n_lin];
4518 let dt = vec![1.0f32; n_lin];
4519 let outp = vec![0.02f32; hidden * value_dim];
4520 add_raw(
4521 "model.layers.0.linear_attn.in_proj_qkvz.weight",
4522 vec![2 * key_dim + 2 * value_dim, hidden],
4523 &qkvz,
4524 );
4525 add_raw(
4526 "model.layers.0.linear_attn.in_proj_ba.weight",
4527 vec![2 * n_lin, hidden],
4528 &ba,
4529 );
4530 add_raw(
4531 "model.layers.0.linear_attn.conv1d.weight",
4532 vec![conv_dim, conv_k],
4533 &conv,
4534 );
4535 add_raw("model.layers.0.linear_attn.A_log", vec![n_lin], &a_log);
4536 add_raw("model.layers.0.linear_attn.dt_bias", vec![n_lin], &dt);
4537 add_raw(
4538 "model.layers.0.linear_attn.out_proj.weight",
4539 vec![hidden, value_dim],
4540 &outp,
4541 );
4542 let g = vec![0.02f32; inter * hidden];
4543 let d = vec![0.02f32; hidden * inter];
4544 add_raw(
4545 "model.layers.0.mlp.gate_proj.weight",
4546 vec![inter, hidden],
4547 &g,
4548 );
4549 add_raw("model.layers.0.mlp.up_proj.weight", vec![inter, hidden], &g);
4550 add_raw(
4551 "model.layers.0.mlp.down_proj.weight",
4552 vec![hidden, inter],
4553 &d,
4554 );
4555
4556 add_raw("model.layers.1.input_layernorm.weight", vec![hidden], &n1);
4557 add_raw(
4558 "model.layers.1.post_attention_layernorm.weight",
4559 vec![hidden],
4560 &n1,
4561 );
4562 let wq = vec![0.02f32; q_dim * hidden];
4563 let wk = vec![0.02f32; k_dim * hidden];
4564 let wv = vec![0.02f32; k_dim * hidden];
4565 let wo = vec![0.02f32; hidden * q_dim];
4566 add_raw(
4567 "model.layers.1.self_attn.q_proj.weight",
4568 vec![q_dim, hidden],
4569 &wq,
4570 );
4571 add_raw(
4572 "model.layers.1.self_attn.k_proj.weight",
4573 vec![k_dim, hidden],
4574 &wk,
4575 );
4576 add_raw(
4577 "model.layers.1.self_attn.v_proj.weight",
4578 vec![k_dim, hidden],
4579 &wv,
4580 );
4581 add_raw(
4582 "model.layers.1.self_attn.o_proj.weight",
4583 vec![hidden, q_dim],
4584 &wo,
4585 );
4586 add_raw(
4587 "model.layers.1.mlp.gate_proj.weight",
4588 vec![inter, hidden],
4589 &g,
4590 );
4591 add_raw("model.layers.1.mlp.up_proj.weight", vec![inter, hidden], &g);
4592 add_raw(
4593 "model.layers.1.mlp.down_proj.weight",
4594 vec![hidden, inter],
4595 &d,
4596 );
4597 add_raw("model.norm.weight", vec![hidden], &n1);
4598 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
4599
4600 let cfg = json!({
4601 "format": "aria-quant-bundle",
4602 "format_version": 2,
4603 "quantization": "test",
4604 "hadamard_seed": 0,
4605 "model": {
4606 "hidden_size": hidden,
4607 "num_layers": 2,
4608 "num_attention_heads": n_heads,
4609 "num_kv_heads": n_kv,
4610 "head_dim": head_dim,
4611 "intermediate_size": inter,
4612 "vocab_size": vocab,
4613 "context_length": 32,
4614 "rope_theta": 10000.0,
4615 "layer_types": ["linear_attention", "full_attention"]
4616 },
4617 "tensors": tensors
4618 });
4619 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
4620 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
4621 let mut s = SessionBuilder::new()
4622 .model(dir.path())
4623 .family("qwen/qwen3.5-2b")
4624 .build()
4625 .unwrap();
4626 assert_eq!(s.config().partial_rotary_factor, Some(0.25));
4627 assert!(
4628 (s.config().rope_theta - 10_000_000.0).abs() < 1.0,
4629 "Qwen3.5 Llama-default rope_theta must become 1e7, got {}",
4630 s.config().rope_theta
4631 );
4632 let gen = s
4633 .generate(
4634 &[1, 2, 3],
4635 &GenerateOpts {
4636 max_tokens: 2,
4637 temperature: 0.0,
4638 },
4639 )
4640 .unwrap();
4641 assert_eq!(gen.tokens.len(), 2);
4642 }
4643
4644 #[test]
4645 fn gated_deltanet_split_qwen35_projections_generate() {
4646 let dir = tempfile::tempdir().unwrap();
4647 let hidden = 8usize;
4648 let inter = 16usize;
4649 let vocab = 16usize;
4650 let n_heads = 2usize;
4651 let n_kv = 1usize;
4652 let head_dim = 4usize;
4653 let q_dim = n_heads * head_dim;
4654 let k_dim = n_kv * head_dim;
4655 let n_lin = 2usize;
4656 let hk = 4usize;
4657 let hv = 4usize;
4658 let key_dim = n_lin * hk;
4659 let value_dim = n_lin * hv;
4660 let conv_k = 4usize;
4661 let conv_dim = key_dim * 2 + value_dim;
4662
4663 let mut tensors = serde_json::Map::new();
4664 let mut bin = Vec::new();
4665 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
4666 let offset = bin.len();
4667 for &v in data {
4668 bin.extend_from_slice(&v.to_le_bytes());
4669 }
4670 let nbytes = data.len() * 4;
4671 let mut meta = serde_json::Map::new();
4672 meta.insert("kind".into(), json!("raw"));
4673 meta.insert("dtype".into(), json!("f32"));
4674 meta.insert("shape".into(), json!(shape));
4675 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
4676 tensors.insert(name.to_string(), Value::Object(meta));
4677 };
4678 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
4679 add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
4680 let n1 = vec![1.0f32; hidden];
4681 add_raw("model.layers.0.input_layernorm.weight", vec![hidden], &n1);
4682 add_raw(
4683 "model.layers.0.post_attention_layernorm.weight",
4684 vec![hidden],
4685 &n1,
4686 );
4687 let qkv = vec![0.02f32; (2 * key_dim + value_dim) * hidden];
4688 let z = vec![0.02f32; value_dim * hidden];
4689 let proj_b = vec![0.1f32; n_lin * hidden];
4690 let proj_a = vec![0.1f32; n_lin * hidden];
4691 let conv = vec![0.05f32; conv_dim * conv_k];
4692 let a_log = vec![0.5f32; n_lin];
4693 let dt = vec![1.0f32; n_lin];
4694 let outp = vec![0.02f32; hidden * value_dim];
4695 add_raw(
4696 "model.layers.0.linear_attn.in_proj_qkv.weight",
4697 vec![2 * key_dim + value_dim, hidden],
4698 &qkv,
4699 );
4700 add_raw(
4701 "model.layers.0.linear_attn.in_proj_z.weight",
4702 vec![value_dim, hidden],
4703 &z,
4704 );
4705 add_raw(
4706 "model.layers.0.linear_attn.in_proj_b.weight",
4707 vec![n_lin, hidden],
4708 &proj_b,
4709 );
4710 add_raw(
4711 "model.layers.0.linear_attn.in_proj_a.weight",
4712 vec![n_lin, hidden],
4713 &proj_a,
4714 );
4715 add_raw(
4716 "model.layers.0.linear_attn.conv1d.weight",
4717 vec![conv_dim, conv_k],
4718 &conv,
4719 );
4720 add_raw("model.layers.0.linear_attn.A_log", vec![n_lin], &a_log);
4721 add_raw("model.layers.0.linear_attn.dt_bias", vec![n_lin], &dt);
4722 add_raw(
4723 "model.layers.0.linear_attn.out_proj.weight",
4724 vec![hidden, value_dim],
4725 &outp,
4726 );
4727 let g = vec![0.02f32; inter * hidden];
4728 let d = vec![0.02f32; hidden * inter];
4729 add_raw(
4730 "model.layers.0.mlp.gate_proj.weight",
4731 vec![inter, hidden],
4732 &g,
4733 );
4734 add_raw("model.layers.0.mlp.up_proj.weight", vec![inter, hidden], &g);
4735 add_raw(
4736 "model.layers.0.mlp.down_proj.weight",
4737 vec![hidden, inter],
4738 &d,
4739 );
4740
4741 add_raw("model.layers.1.input_layernorm.weight", vec![hidden], &n1);
4742 add_raw(
4743 "model.layers.1.post_attention_layernorm.weight",
4744 vec![hidden],
4745 &n1,
4746 );
4747 let wq = vec![0.02f32; q_dim * hidden];
4748 let wk = vec![0.02f32; k_dim * hidden];
4749 let wv = vec![0.02f32; k_dim * hidden];
4750 let wo = vec![0.02f32; hidden * q_dim];
4751 add_raw(
4752 "model.layers.1.self_attn.q_proj.weight",
4753 vec![q_dim, hidden],
4754 &wq,
4755 );
4756 add_raw(
4757 "model.layers.1.self_attn.k_proj.weight",
4758 vec![k_dim, hidden],
4759 &wk,
4760 );
4761 add_raw(
4762 "model.layers.1.self_attn.v_proj.weight",
4763 vec![k_dim, hidden],
4764 &wv,
4765 );
4766 add_raw(
4767 "model.layers.1.self_attn.o_proj.weight",
4768 vec![hidden, q_dim],
4769 &wo,
4770 );
4771 add_raw(
4772 "model.layers.1.mlp.gate_proj.weight",
4773 vec![inter, hidden],
4774 &g,
4775 );
4776 add_raw("model.layers.1.mlp.up_proj.weight", vec![inter, hidden], &g);
4777 add_raw(
4778 "model.layers.1.mlp.down_proj.weight",
4779 vec![hidden, inter],
4780 &d,
4781 );
4782 add_raw("model.norm.weight", vec![hidden], &n1);
4783 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
4784
4785 let cfg = json!({
4786 "format": "aria-quant-bundle",
4787 "format_version": 2,
4788 "quantization": "test",
4789 "hadamard_seed": 0,
4790 "model": {
4791 "hidden_size": hidden,
4792 "num_layers": 2,
4793 "num_attention_heads": n_heads,
4794 "num_kv_heads": n_kv,
4795 "head_dim": head_dim,
4796 "intermediate_size": inter,
4797 "vocab_size": vocab,
4798 "context_length": 32,
4799 "rope_theta": 10000.0,
4800 "layer_types": ["linear_attention", "full_attention"]
4801 },
4802 "tensors": tensors
4803 });
4804 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
4805 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
4806 let mut s = SessionBuilder::new()
4807 .model(dir.path())
4808 .family("qwen/qwen3.5-0.8b")
4809 .build()
4810 .unwrap();
4811 let gen = s
4812 .generate(
4813 &[1, 2, 3],
4814 &GenerateOpts {
4815 max_tokens: 2,
4816 temperature: 0.0,
4817 },
4818 )
4819 .unwrap();
4820 assert_eq!(gen.tokens.len(), 2);
4821 }
4822
4823 #[test]
4824 fn vision_and_action_consume_bundle_weights() {
4825 let dir = tempfile::tempdir().unwrap();
4826 write_tiny_q4_bundle(dir.path()).unwrap();
4827 let cfg_path = dir.path().join("config.json");
4828 let mut cfg: Value =
4829 serde_json::from_str(&std::fs::read_to_string(&cfg_path).unwrap()).unwrap();
4830 let hidden = cfg["model"]["hidden_size"].as_u64().unwrap() as usize;
4831 let mut tensors = cfg["tensors"].as_object().cloned().unwrap();
4832 let mut bin = std::fs::read(dir.path().join("weight.bin")).unwrap();
4833 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
4834 let offset = bin.len();
4835 for &v in data {
4836 bin.extend_from_slice(&v.to_le_bytes());
4837 }
4838 let nbytes = data.len() * 4;
4839 tensors.insert(
4840 name.to_string(),
4841 json!({
4842 "kind": "raw",
4843 "dtype": "f32",
4844 "shape": shape,
4845 "offsets": { "data": [offset, nbytes] }
4846 }),
4847 );
4848 };
4849 let vis = vec![0.1f32; hidden * 3];
4850 add_raw("mm_projector.weight", vec![hidden, 3], &vis);
4851 let act_dim = 7usize;
4852 let act = vec![0.05f32; act_dim * hidden];
4853 add_raw("action_head.weight", vec![act_dim, hidden], &act);
4854 cfg["tensors"] = Value::Object(tensors);
4855 std::fs::write(&cfg_path, cfg.to_string()).unwrap();
4856 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
4857
4858 let s = SessionBuilder::new()
4859 .model(dir.path())
4860 .family("lfm/lfm2-vl-450m")
4861 .build()
4862 .unwrap();
4863 let rgb = vec![10u8; 3 * 4 * 4];
4864 let pref = s.vision_prefix(&rgb, 4, 4).unwrap();
4865 assert_eq!(pref.len(), hidden);
4866
4867 let vla = SessionBuilder::new()
4868 .model(dir.path())
4869 .family("openvla/openvla-7b")
4870 .build()
4871 .unwrap();
4872 let a = vla.predict_action("move", act_dim).unwrap();
4873 assert_eq!(a.len(), act_dim);
4874 }
4875
4876 #[test]
4877 fn load_real_hf_named_bundle_if_present() {
4878 let Ok(path) = std::env::var("ARIA_SMOKE_BUNDLE") else {
4880 return;
4881 };
4882 let path = std::path::Path::new(&path);
4883 if !path.join("config.json").is_file() {
4884 return;
4885 }
4886 let family = if path.to_string_lossy().contains("gemma-4") {
4887 "gemma/gemma-4-e2b-it"
4888 } else {
4889 "qwen/qwen3-0.6b"
4890 };
4891 let s = SessionBuilder::new()
4892 .model(path)
4893 .family(family)
4894 .build()
4895 .unwrap_or_else(|e| panic!("{family} bundle should materialize: {e}"));
4896 assert!(s.config().num_layers > 0);
4897 assert!(s.config().hidden_size > 0);
4898 if family.contains("gemma-4") && s.config().hidden_size >= 1024 {
4899 assert!(
4900 s.weights.ple.is_some(),
4901 "real Gemma-4 q4 must load codebook PLE"
4902 );
4903 let hidden = s.config().hidden_size;
4904 let vocab = s.config().vocab_size;
4905 assert!(
4906 s.weights.emb.data.len() >= vocab.saturating_mul(hidden),
4907 "embed table too small for vocab={vocab} hidden={hidden}"
4908 );
4909 }
4910 }
4911
4912 #[test]
4913 fn gemma4_four_norm_ple_and_tied_embed_generate() {
4914 let dir = tempfile::tempdir().unwrap();
4915 let hidden = 8usize;
4916 let layers = 1usize;
4917 let inter = 16usize;
4918 let vocab = 16usize;
4919 let n_heads = 2usize;
4920 let n_kv = 1usize;
4921 let head_dim = 4usize;
4922 let q_dim = n_heads * head_dim;
4923 let k_dim = n_kv * head_dim;
4924 let ple_d = 4usize;
4925 let p = "model.language_model";
4926
4927 let mut tensors = serde_json::Map::new();
4928 let mut bin = Vec::new();
4929 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
4930 let offset = bin.len();
4931 for &v in data {
4932 bin.extend_from_slice(&v.to_le_bytes());
4933 }
4934 let nbytes = data.len() * 4;
4935 let mut meta = serde_json::Map::new();
4936 meta.insert("kind".into(), json!("raw"));
4937 meta.insert("dtype".into(), json!("f32"));
4938 meta.insert("shape".into(), json!(shape));
4939 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
4940 tensors.insert(name.to_string(), Value::Object(meta));
4941 };
4942 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
4943 add_raw(
4944 &format!("{p}.embed_tokens.weight"),
4945 vec![vocab, hidden],
4946 &emb,
4947 );
4948 let n1 = vec![1.0f32; hidden];
4949 add_raw(
4950 &format!("{p}.layers.0.input_layernorm.weight"),
4951 vec![hidden],
4952 &n1,
4953 );
4954 add_raw(
4955 &format!("{p}.layers.0.post_attention_layernorm.weight"),
4956 vec![hidden],
4957 &n1,
4958 );
4959 add_raw(
4960 &format!("{p}.layers.0.pre_feedforward_layernorm.weight"),
4961 vec![hidden],
4962 &n1,
4963 );
4964 add_raw(
4965 &format!("{p}.layers.0.post_feedforward_layernorm.weight"),
4966 vec![hidden],
4967 &n1,
4968 );
4969 add_raw(&format!("{p}.layers.0.layer_scalar"), vec![1], &[0.5f32]);
4971 let wq = vec![0.01f32; q_dim * hidden];
4972 let wk = vec![0.01f32; k_dim * hidden];
4973 let wv = vec![0.01f32; k_dim * hidden];
4974 let wo = vec![0.01f32; hidden * q_dim];
4975 add_raw(
4976 &format!("{p}.layers.0.self_attn.q_proj.weight"),
4977 vec![q_dim, hidden],
4978 &wq,
4979 );
4980 add_raw(
4981 &format!("{p}.layers.0.self_attn.k_proj.weight"),
4982 vec![k_dim, hidden],
4983 &wk,
4984 );
4985 add_raw(
4986 &format!("{p}.layers.0.self_attn.v_proj.weight"),
4987 vec![k_dim, hidden],
4988 &wv,
4989 );
4990 add_raw(
4991 &format!("{p}.layers.0.self_attn.o_proj.weight"),
4992 vec![hidden, q_dim],
4993 &wo,
4994 );
4995 let g = vec![0.01f32; inter * hidden];
4996 let d = vec![0.01f32; hidden * inter];
4997 add_raw(
4998 &format!("{p}.layers.0.mlp.gate_proj.weight"),
4999 vec![inter, hidden],
5000 &g,
5001 );
5002 add_raw(
5003 &format!("{p}.layers.0.mlp.up_proj.weight"),
5004 vec![inter, hidden],
5005 &g,
5006 );
5007 add_raw(
5008 &format!("{p}.layers.0.mlp.down_proj.weight"),
5009 vec![hidden, inter],
5010 &d,
5011 );
5012 let packed = layers * ple_d;
5013 let ple_emb = vec![0.02f32; vocab * packed];
5014 add_raw(
5015 &format!("{p}.embed_tokens_per_layer.weight"),
5016 vec![vocab, packed],
5017 &ple_emb,
5018 );
5019 let ple_proj = vec![0.01f32; packed * hidden];
5020 add_raw(
5021 &format!("{p}.per_layer_model_projection.weight"),
5022 vec![packed, hidden],
5023 &ple_proj,
5024 );
5025 let ple_pn = vec![1.0f32; ple_d];
5026 add_raw(
5027 &format!("{p}.per_layer_projection_norm.weight"),
5028 vec![ple_d],
5029 &ple_pn,
5030 );
5031 let ple_gate = vec![0.01f32; ple_d * hidden];
5032 let ple_out = vec![0.01f32; hidden * ple_d];
5033 add_raw(
5034 &format!("{p}.layers.0.per_layer_input_gate.weight"),
5035 vec![ple_d, hidden],
5036 &ple_gate,
5037 );
5038 add_raw(
5039 &format!("{p}.layers.0.per_layer_projection.weight"),
5040 vec![hidden, ple_d],
5041 &ple_out,
5042 );
5043 add_raw(
5044 &format!("{p}.layers.0.post_per_layer_input_norm.weight"),
5045 vec![hidden],
5046 &n1,
5047 );
5048 add_raw(&format!("{p}.norm.weight"), vec![hidden], &n1);
5049 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
5050
5051 let cfg = json!({
5052 "format": "aria-quant-bundle",
5053 "format_version": 2,
5054 "quantization": "test",
5055 "group_size_default": 32,
5056 "hadamard_seed": 0,
5057 "model": {
5058 "hidden_size": hidden,
5059 "num_layers": layers,
5060 "num_attention_heads": n_heads,
5061 "num_kv_heads": n_kv,
5062 "intermediate_size": inter,
5063 "vocab_size": vocab,
5064 "context_length": 32,
5065 "rope_theta": 10000.0,
5066 "hidden_act": "gelu_pytorch_tanh",
5067 "tie_word_embeddings": true,
5068 "head_dim": head_dim,
5069 "global_head_dim": head_dim,
5070 "sliding_window": 512,
5071 "partial_rotary_factor": 0.25,
5072 "layer_types": ["full_attention"]
5073 },
5074 "tensors": tensors
5075 });
5076 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
5077 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
5078
5079 let mut s = SessionBuilder::new()
5080 .model(dir.path())
5081 .family("gemma/gemma-4-e2b-it")
5082 .build()
5083 .unwrap();
5084 assert!((s.embed_scale - (hidden as f32).sqrt()).abs() < 1e-5);
5085 assert!(s.weights.ple.is_some());
5086 assert!((s.weights.layers[0].layer_scalar - 0.5).abs() < 1e-6);
5087 assert!(s.weights.layers[0].post_attn_norm.is_some());
5088 assert!(s.weights.layers[0].post_ffn_norm.is_some());
5089 let prompt = vec![1u32, 2];
5090 let batched = s
5091 .generate(
5092 &prompt,
5093 &GenerateOpts {
5094 max_tokens: 3,
5095 temperature: 0.0,
5096 },
5097 )
5098 .unwrap();
5099 let step = s
5100 .generate(
5101 &prompt,
5102 &GenerateOpts {
5103 max_tokens: 3,
5104 temperature: 0.0,
5105 },
5106 )
5107 .unwrap();
5108 assert_eq!(batched.tokens, step.tokens);
5109 assert_eq!(batched.tokens.len(), 3);
5110 assert_eq!(s.config().sliding_window, Some(512));
5111 }
5112
5113 #[test]
5114 fn gemma4_ple_required_gate() {
5115 assert!(!gemma4_requires_ple("gemma/gemma-4-e2b-it", 64));
5116 assert!(gemma4_requires_ple("gemma/gemma-4-e2b-it", 1024));
5117 assert!(gemma4_requires_ple("gemma/gemma-4-e2b-it", 1536));
5118 assert!(!gemma4_requires_ple("gemma/gemma-3-1b-it", 1152));
5119 assert!(gemma4_requires_ple("gemma/gemma-3n-e2b-it", 2048));
5120 assert!(!gemma4_requires_ple("qwen/qwen3-0.6b", 1536));
5121 assert!(!gemma3n_requires_altup("gemma/gemma-3n-e2b-it", 64));
5122 assert!(gemma3n_requires_altup("gemma/gemma-3n-e2b-it", 2048));
5123 assert!(!gemma3n_requires_altup("gemma/gemma-4-e2b-it", 2048));
5124 }
5125
5126 #[test]
5127 fn gemma3n_altup_laurel_ple_generate() {
5128 let dir = tempfile::tempdir().unwrap();
5129 let hidden = 8usize;
5130 let layers = 1usize;
5131 let inter = 16usize;
5132 let vocab = 16usize;
5133 let n_heads = 2usize;
5134 let n_kv = 1usize;
5135 let head_dim = 4usize;
5136 let q_dim = n_heads * head_dim;
5137 let k_dim = n_kv * head_dim;
5138 let ple_d = 4usize;
5139 let rank = 2usize;
5140 let n_alt = 4usize;
5141 let p = "model.language_model";
5142
5143 let mut tensors = serde_json::Map::new();
5144 let mut bin = Vec::new();
5145 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
5146 let offset = bin.len();
5147 for &v in data {
5148 bin.extend_from_slice(&v.to_le_bytes());
5149 }
5150 let nbytes = data.len() * 4;
5151 let mut meta = serde_json::Map::new();
5152 meta.insert("kind".into(), json!("raw"));
5153 meta.insert("dtype".into(), json!("f32"));
5154 meta.insert("shape".into(), json!(shape));
5155 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
5156 tensors.insert(name.to_string(), Value::Object(meta));
5157 };
5158 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
5159 add_raw(
5160 &format!("{p}.embed_tokens.weight"),
5161 vec![vocab, hidden],
5162 &emb,
5163 );
5164 let n1 = vec![1.0f32; hidden];
5165 add_raw(
5166 &format!("{p}.layers.0.input_layernorm.weight"),
5167 vec![hidden],
5168 &n1,
5169 );
5170 add_raw(
5171 &format!("{p}.layers.0.post_attention_layernorm.weight"),
5172 vec![hidden],
5173 &n1,
5174 );
5175 add_raw(
5176 &format!("{p}.layers.0.pre_feedforward_layernorm.weight"),
5177 vec![hidden],
5178 &n1,
5179 );
5180 add_raw(
5181 &format!("{p}.layers.0.post_feedforward_layernorm.weight"),
5182 vec![hidden],
5183 &n1,
5184 );
5185 let wq = vec![0.01f32; q_dim * hidden];
5186 let wk = vec![0.01f32; k_dim * hidden];
5187 let wv = vec![0.01f32; k_dim * hidden];
5188 let wo = vec![0.01f32; hidden * q_dim];
5189 add_raw(
5190 &format!("{p}.layers.0.self_attn.q_proj.weight"),
5191 vec![q_dim, hidden],
5192 &wq,
5193 );
5194 add_raw(
5195 &format!("{p}.layers.0.self_attn.k_proj.weight"),
5196 vec![k_dim, hidden],
5197 &wk,
5198 );
5199 add_raw(
5200 &format!("{p}.layers.0.self_attn.v_proj.weight"),
5201 vec![k_dim, hidden],
5202 &wv,
5203 );
5204 add_raw(
5205 &format!("{p}.layers.0.self_attn.o_proj.weight"),
5206 vec![hidden, q_dim],
5207 &wo,
5208 );
5209 let g = vec![0.01f32; inter * hidden];
5210 let d = vec![0.01f32; hidden * inter];
5211 add_raw(
5212 &format!("{p}.layers.0.mlp.gate_proj.weight"),
5213 vec![inter, hidden],
5214 &g,
5215 );
5216 add_raw(
5217 &format!("{p}.layers.0.mlp.up_proj.weight"),
5218 vec![inter, hidden],
5219 &g,
5220 );
5221 add_raw(
5222 &format!("{p}.layers.0.mlp.down_proj.weight"),
5223 vec![hidden, inter],
5224 &d,
5225 );
5226 let packed = layers * ple_d;
5227 add_raw(
5228 &format!("{p}.embed_tokens_per_layer.weight"),
5229 vec![vocab, packed],
5230 &vec![0.02f32; vocab * packed],
5231 );
5232 add_raw(
5233 &format!("{p}.per_layer_model_projection.weight"),
5234 vec![packed, hidden],
5235 &vec![0.01f32; packed * hidden],
5236 );
5237 add_raw(
5238 &format!("{p}.per_layer_projection_norm.weight"),
5239 vec![ple_d],
5240 &vec![1.0f32; ple_d],
5241 );
5242 add_raw(
5243 &format!("{p}.layers.0.per_layer_input_gate.weight"),
5244 vec![ple_d, hidden],
5245 &vec![0.01f32; ple_d * hidden],
5246 );
5247 add_raw(
5248 &format!("{p}.layers.0.per_layer_projection.weight"),
5249 vec![hidden, ple_d],
5250 &vec![0.01f32; hidden * ple_d],
5251 );
5252 add_raw(
5253 &format!("{p}.layers.0.post_per_layer_input_norm.weight"),
5254 vec![hidden],
5255 &n1,
5256 );
5257 add_raw(
5258 &format!("{p}.layers.0.altup.modality_router.weight"),
5259 vec![n_alt, hidden],
5260 &vec![0.01f32; n_alt * hidden],
5261 );
5262 add_raw(
5263 &format!("{p}.layers.0.altup.router_norm.weight"),
5264 vec![hidden],
5265 &n1,
5266 );
5267 add_raw(
5268 &format!("{p}.layers.0.altup.prediction_coefs.weight"),
5269 vec![n_alt * n_alt, n_alt],
5270 &vec![0.0f32; n_alt * n_alt * n_alt],
5271 );
5272 add_raw(
5273 &format!("{p}.layers.0.altup.correction_coefs.weight"),
5274 vec![n_alt, n_alt],
5275 &vec![0.0f32; n_alt * n_alt],
5276 );
5277 add_raw(
5278 &format!("{p}.layers.0.altup.correct_output_scale"),
5279 vec![hidden],
5280 &n1,
5281 );
5282 add_raw(
5283 &format!("{p}.layers.0.laurel.linear_left.weight"),
5284 vec![rank, hidden],
5285 &vec![0.01f32; rank * hidden],
5286 );
5287 add_raw(
5288 &format!("{p}.layers.0.laurel.linear_right.weight"),
5289 vec![hidden, rank],
5290 &vec![0.01f32; hidden * rank],
5291 );
5292 add_raw(
5293 &format!("{p}.layers.0.laurel.post_laurel_norm.weight"),
5294 vec![hidden],
5295 &n1,
5296 );
5297 let eye: Vec<f32> = (0..hidden * hidden)
5298 .map(|i| if i / hidden == i % hidden { 0.05 } else { 0.0 })
5299 .collect();
5300 for i in 0..3 {
5301 add_raw(
5302 &format!("{p}.altup_projections.{i}.weight"),
5303 vec![hidden, hidden],
5304 &eye,
5305 );
5306 add_raw(
5307 &format!("{p}.altup_unembed_projections.{i}.weight"),
5308 vec![hidden, hidden],
5309 &eye,
5310 );
5311 }
5312 add_raw(&format!("{p}.norm.weight"), vec![hidden], &n1);
5313 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
5314
5315 let cfg = json!({
5316 "format": "aria-quant-bundle",
5317 "format_version": 2,
5318 "quantization": "test",
5319 "group_size_default": 32,
5320 "hadamard_seed": 0,
5321 "model": {
5322 "hidden_size": hidden,
5323 "num_layers": layers,
5324 "num_attention_heads": n_heads,
5325 "num_kv_heads": n_kv,
5326 "intermediate_size": inter,
5327 "vocab_size": vocab,
5328 "context_length": 32,
5329 "rope_theta": 1000000.0,
5330 "hidden_act": "gelu_pytorch_tanh",
5331 "tie_word_embeddings": true,
5332 "head_dim": head_dim,
5333 "sliding_window": 512,
5334 "layer_types": ["full_attention"]
5335 },
5336 "tensors": tensors
5337 });
5338 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
5339 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
5340
5341 let mut s = SessionBuilder::new()
5342 .model(dir.path())
5343 .family("gemma/gemma-3n-e2b-it")
5344 .build()
5345 .unwrap();
5346 assert!(s.has_gemma3n_graph());
5347 assert!(s.weights.ple.is_some());
5348 assert!(s.weights.layers[0].altup.is_some());
5349 assert!(s.weights.layers[0].laurel.is_some());
5350 let prompt = vec![1u32, 2];
5351 let batched = s
5352 .generate(
5353 &prompt,
5354 &GenerateOpts {
5355 max_tokens: 3,
5356 temperature: 0.0,
5357 },
5358 )
5359 .unwrap();
5360 let step = s
5361 .generate(
5362 &prompt,
5363 &GenerateOpts {
5364 max_tokens: 3,
5365 temperature: 0.0,
5366 },
5367 )
5368 .unwrap();
5369 assert_eq!(batched.tokens, step.tokens);
5370 assert_eq!(batched.tokens.len(), 3);
5371 }
5372
5373 #[test]
5374 fn gaussian_topk_sparsity_zeros_below_cutoff() {
5375 let row: Vec<f32> = (0..100).map(|i| i as f32).collect();
5376 let y = gaussian_topk(&row, 100, 0.95).unwrap();
5377 assert!(y.iter().all(|&v| v >= 0.0));
5378 assert!(y.iter().filter(|&&v| v == 0.0).count() > 50);
5379 assert!(y.iter().any(|&v| v > 0.0));
5380 }
5381
5382 #[test]
5383 fn gemma4_e2b_scale_missing_ple_is_hard_error() {
5384 let dir = tempfile::tempdir().unwrap();
5385 let hidden = 1024usize;
5386 let layers = 1usize;
5387 let vocab = 8usize;
5388 let n_heads = 8usize;
5389 let n_kv = 1usize;
5390 let head_dim = 128usize;
5391 let q_dim = n_heads * head_dim;
5392 let k_dim = n_kv * head_dim;
5393 let inter = 32usize;
5394 let p = "model.language_model";
5395
5396 let mut tensors = serde_json::Map::new();
5397 let mut bin = Vec::new();
5398 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
5399 let offset = bin.len();
5400 for &v in data {
5401 bin.extend_from_slice(&v.to_le_bytes());
5402 }
5403 let nbytes = data.len() * 4;
5404 tensors.insert(
5405 name.to_string(),
5406 json!({
5407 "kind": "raw",
5408 "dtype": "f32",
5409 "shape": shape,
5410 "offsets": { "data": [offset, nbytes] }
5411 }),
5412 );
5413 };
5414 let emb = vec![0.01f32; vocab * hidden];
5415 add_raw(
5416 &format!("{p}.embed_tokens.weight"),
5417 vec![vocab, hidden],
5418 &emb,
5419 );
5420 let ones = vec![1.0f32; hidden];
5421 add_raw(
5422 &format!("{p}.layers.0.input_layernorm.weight"),
5423 vec![hidden],
5424 &ones,
5425 );
5426 add_raw(
5427 &format!("{p}.layers.0.pre_feedforward_layernorm.weight"),
5428 vec![hidden],
5429 &ones,
5430 );
5431 add_raw(
5432 &format!("{p}.layers.0.post_attention_layernorm.weight"),
5433 vec![hidden],
5434 &ones,
5435 );
5436 add_raw(
5437 &format!("{p}.layers.0.post_feedforward_layernorm.weight"),
5438 vec![hidden],
5439 &ones,
5440 );
5441 add_raw(&format!("{p}.norm.weight"), vec![hidden], &ones);
5442 let q = vec![0.01f32; q_dim * hidden];
5443 let k = vec![0.01f32; k_dim * hidden];
5444 add_raw(
5445 &format!("{p}.layers.0.self_attn.q_proj.weight"),
5446 vec![q_dim, hidden],
5447 &q,
5448 );
5449 add_raw(
5450 &format!("{p}.layers.0.self_attn.k_proj.weight"),
5451 vec![k_dim, hidden],
5452 &k,
5453 );
5454 add_raw(
5455 &format!("{p}.layers.0.self_attn.v_proj.weight"),
5456 vec![k_dim, hidden],
5457 &k,
5458 );
5459 add_raw(
5460 &format!("{p}.layers.0.self_attn.o_proj.weight"),
5461 vec![hidden, q_dim],
5462 &q,
5463 );
5464 let qn = vec![1.0f32; head_dim];
5465 add_raw(
5466 &format!("{p}.layers.0.self_attn.q_norm.weight"),
5467 vec![head_dim],
5468 &qn,
5469 );
5470 add_raw(
5471 &format!("{p}.layers.0.self_attn.k_norm.weight"),
5472 vec![head_dim],
5473 &qn,
5474 );
5475 let g = vec![0.01f32; inter * hidden];
5476 add_raw(
5477 &format!("{p}.layers.0.mlp.gate_proj.weight"),
5478 vec![inter, hidden],
5479 &g,
5480 );
5481 add_raw(
5482 &format!("{p}.layers.0.mlp.up_proj.weight"),
5483 vec![inter, hidden],
5484 &g,
5485 );
5486 add_raw(
5487 &format!("{p}.layers.0.mlp.down_proj.weight"),
5488 vec![hidden, inter],
5489 &g,
5490 );
5491 let cfg = json!({
5492 "format": "aria-quant-bundle",
5493 "format_version": 2,
5494 "quantization": "test",
5495 "group_size_default": 32,
5496 "hadamard_seed": 0,
5497 "model": {
5498 "hidden_size": hidden,
5499 "num_layers": layers,
5500 "num_attention_heads": n_heads,
5501 "num_kv_heads": n_kv,
5502 "intermediate_size": inter,
5503 "vocab_size": vocab,
5504 "context_length": 32,
5505 "rope_theta": 10000.0,
5506 "hidden_act": "gelu_pytorch_tanh",
5507 "tie_word_embeddings": true,
5508 "head_dim": head_dim,
5509 "global_head_dim": head_dim,
5510 "sliding_window": 512,
5511 "partial_rotary_factor": 0.25,
5512 "num_kv_shared_layers": 0,
5513 "layer_types": ["full_attention"]
5514 },
5515 "tensors": tensors
5516 });
5517 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
5518 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
5519
5520 let err = SessionBuilder::new()
5521 .model(dir.path())
5522 .family("gemma/gemma-4-e2b-it")
5523 .build()
5524 .unwrap_err();
5525 let msg = err.to_string();
5526 assert!(
5527 msg.contains("PLE") && msg.contains("embed_tokens_per_layer"),
5528 "{msg}"
5529 );
5530
5531 let named_parent = tempfile::tempdir().unwrap();
5535 let named = named_parent.path().join("gemma-3-1b-it_q326");
5536 std::fs::create_dir(&named).unwrap();
5537 std::fs::copy(dir.path().join("config.json"), named.join("config.json")).unwrap();
5538 std::fs::copy(dir.path().join("weight.bin"), named.join("weight.bin")).unwrap();
5539 let s = SessionBuilder::new().model(&named).build().unwrap();
5540 assert_eq!(s.family().path(), "gemma/gemma-3-1b-it");
5541 }
5542
5543 #[test]
5544 fn gemma4_sliding_window_config_and_generate() {
5545 let dir_wide = tempfile::tempdir().unwrap();
5546 write_tiny_q4_bundle(dir_wide.path()).unwrap();
5547 let dir_narrow = tempfile::tempdir().unwrap();
5548 write_tiny_q4_bundle(dir_narrow.path()).unwrap();
5549 let patch = |path: &std::path::Path, window: usize| {
5550 let cfg_path = path.join("config.json");
5551 let raw = std::fs::read_to_string(&cfg_path).unwrap();
5552 let mut cfg: Value = serde_json::from_str(&raw).unwrap();
5553 cfg["model"]["sliding_window"] = json!(window);
5554 cfg["model"]["layer_types"] = json!(["sliding_attention", "sliding_attention"]);
5555 std::fs::write(&cfg_path, cfg.to_string()).unwrap();
5556 };
5557 patch(dir_wide.path(), 512);
5558 patch(dir_narrow.path(), 1);
5559 let wide = SessionBuilder::new()
5560 .model(dir_wide.path())
5561 .family("gemma/gemma-4-e2b-it")
5562 .build()
5563 .unwrap();
5564 let mut narrow = SessionBuilder::new()
5565 .model(dir_narrow.path())
5566 .family("gemma/gemma-4-e2b-it")
5567 .build()
5568 .unwrap();
5569 assert_eq!(wide.config().sliding_window, Some(512));
5570 assert_eq!(narrow.config().sliding_window, Some(1));
5571 assert_eq!(wide.attn_window(AttnKind::Sliding), Some(512));
5572 assert_eq!(narrow.attn_window(AttnKind::Sliding), Some(1));
5573 for layer in &narrow.weights.layers {
5574 if let LayerOp::Attn(attn) = &layer.op {
5575 assert_eq!(attn.kind, AttnKind::Sliding);
5576 }
5577 }
5578 let prompt = vec![1u32, 2, 3, 4];
5579 let gen = narrow
5580 .generate(
5581 &prompt,
5582 &GenerateOpts {
5583 max_tokens: 3,
5584 temperature: 0.0,
5585 },
5586 )
5587 .unwrap();
5588 assert_eq!(gen.tokens.len(), 3);
5589
5590 let mut incr = SessionBuilder::new()
5592 .model(dir_narrow.path())
5593 .family("gemma/gemma-4-e2b-it")
5594 .build()
5595 .unwrap();
5596 let again = incr
5597 .generate(
5598 &prompt,
5599 &GenerateOpts {
5600 max_tokens: 3,
5601 temperature: 0.0,
5602 },
5603 )
5604 .unwrap();
5605 assert_eq!(gen.tokens, again.tokens);
5606 }
5607}