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