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 layer: &LayerWeights,
2117 li: usize,
2118 ple_tok: Option<&[f32]>,
2119 hidden: usize,
2120 ) -> Result<(), EngineError> {
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.into_iter()) {
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 layer,
2548 li,
2549 ple_tok.as_deref(),
2550 hidden,
2551 )?;
2552 continue;
2553 }
2554 self.add_normed_residual(&mut x, &ao, layer.post_attn_norm.as_deref())?;
2555 }
2556 LayerOp::Conv(_) | LayerOp::Linear(_) => {
2557 return Err(EngineError::Unsupported(
2558 "batched prefill is only implemented for attention+dense FFN layers".into(),
2559 ));
2560 }
2561 }
2562 let xn2 = self.norm(&x, &layer.ffn_norm)?;
2563 let down = self.apply_ffn(layer, &xn2, hidden)?;
2564 self.add_normed_residual(&mut x, &down, layer.post_ffn_norm.as_deref())?;
2565 self.apply_ple(&mut x, layer, li, ple_tok.as_deref(), hidden)?;
2566 Self::apply_layer_scalar(&mut x, layer.layer_scalar);
2567 }
2568 if let Some(streams) = gemma3n_streams {
2569 x = self.gemma3n_unembed_streams(&streams, hidden)?;
2570 }
2571 state.pos = pos0 + seq;
2572 let last = &x[(seq - 1) * hidden..seq * hidden];
2573 let xn = self.norm(last, &self.weights.output_norm)?;
2574 if !self.weights.output.data.len().is_multiple_of(hidden) {
2575 return Err(EngineError::ShapeMismatch(format!(
2576 "lm_head len {} not divisible by hidden {hidden}",
2577 self.weights.output.data.len()
2578 )));
2579 }
2580 let out_rows = self.weights.output.data.len() / hidden;
2581 let logits = self.wmm(
2582 &self.weights.output,
2583 &xn,
2584 out_rows,
2585 hidden,
2586 GemmAcct::LmHead,
2587 )?;
2588 Ok(self.softcap_logits(logits))
2589 }
2590
2591 fn forward_step(&mut self, tok: u32) -> Result<Vec<f32>, EngineError> {
2592 let mut owned = self.decode.take().ok_or_else(|| {
2593 EngineError::InvalidParam("decode state missing; call generate/prefill first".into())
2594 })?;
2595 let logits = self.forward_step_with(&mut owned, tok);
2596 self.decode = Some(owned);
2597 logits
2598 }
2599
2600 fn forward_step_with(
2601 &self,
2602 state: &mut DecodeState,
2603 tok: u32,
2604 ) -> Result<Vec<f32>, EngineError> {
2605 let hidden = self.conf.hidden_size;
2606 let n_heads = self.conf.num_attention_heads;
2607 let n_kv = self.conf.num_kv_heads;
2608 let vocab = self.conf.vocab_size;
2609 if hidden == 0 {
2610 return Err(EngineError::ShapeMismatch("hidden_size is 0".into()));
2611 }
2612 if self.weights.emb.data.len() < vocab.saturating_mul(hidden)
2613 || !self.weights.emb.data.len().is_multiple_of(hidden)
2614 {
2615 return Err(EngineError::ShapeMismatch(format!(
2616 "embedding length {} not compatible with vocab={vocab} hidden={hidden}",
2617 self.weights.emb.data.len()
2618 )));
2619 }
2620 let pos = state.pos;
2621 let tid = (tok as usize) % vocab;
2622 let mut x = vec![0.0f32; hidden];
2623 x.copy_from_slice(&self.weights.emb.data[tid * hidden..(tid + 1) * hidden]);
2624 if self.embed_scale != 1.0 {
2625 for v in &mut x {
2626 *v *= self.embed_scale;
2627 }
2628 }
2629 let ple_tok = self.compute_ple_inputs(&[tok], &x)?;
2630 let mut gemma3n_streams = if self.has_gemma3n_graph() {
2631 Some(self.gemma3n_expand_streams(&x, hidden)?)
2632 } else {
2633 None
2634 };
2635
2636 for (li, layer) in self.weights.layers.iter().enumerate() {
2637 let mut gemma3n_preds = None;
2638 let mut gemma3n_laurel = None;
2639 let xn = if let Some(ref streams) = gemma3n_streams {
2640 let preds = self.altup_predict(layer, streams, hidden)?;
2641 let xn = self.norm(&preds[0], &layer.attn_norm)?;
2642 gemma3n_laurel = Some(self.apply_laurel(layer, &xn, hidden)?);
2643 gemma3n_preds = Some(preds);
2644 xn
2645 } else {
2646 self.norm(&x, &layer.attn_norm)?
2647 };
2648 match &layer.op {
2649 LayerOp::Attn(attn) => {
2650 if attn.wq.data.len() % hidden != 0 {
2651 return Err(EngineError::ShapeMismatch(
2652 "attn q proj weight not divisible by hidden_size".into(),
2653 ));
2654 }
2655 let q_out = attn.wq.data.len() / hidden;
2656 let q_dim = if attn.q_gate { q_out / 2 } else { q_out };
2657 let head_dim = self.layer_head_dim(attn.kind, q_dim, n_heads)?;
2658 if attn.wo.data.len() != hidden * q_dim {
2659 return Err(EngineError::ShapeMismatch(format!(
2660 "attn output proj weight shape mismatch (wo_len={} hidden={hidden} q_dim={q_dim})",
2661 attn.wo.data.len()
2662 )));
2663 }
2664 let mixed = self.wmm(&attn.wq, &xn, q_out, hidden, GemmAcct::Attn)?;
2665 let (mut q, gate) = if attn.q_gate {
2666 split_interleaved_q_gate(&mixed, 1, n_heads, head_dim)?
2667 } else {
2668 (mixed, Vec::new())
2669 };
2670 if let Some(qn) = &attn.q_norm {
2671 if qn.len() != head_dim {
2672 return Err(EngineError::ShapeMismatch(format!(
2673 "q_norm len {} != head_dim {head_dim}",
2674 qn.len()
2675 )));
2676 }
2677 q = self.norm(&q, qn)?;
2678 }
2679 let (theta, rope) = self.layer_rope_params(attn.kind);
2680 Self::apply_rope(&mut q, head_dim, pos, theta, rope)?;
2681
2682 let (k_src, v_src) = if let (Some(wk), Some(wv)) = (&attn.wk, &attn.wv) {
2683 if wk.data.len() % hidden != 0 || wv.data.len() % hidden != 0 {
2684 return Err(EngineError::ShapeMismatch(
2685 "attn kv proj weight not divisible by hidden_size".into(),
2686 ));
2687 }
2688 let k_dim = wk.data.len() / hidden;
2689 let v_dim = wv.data.len() / hidden;
2690 if k_dim != n_kv * head_dim || v_dim != n_kv * head_dim {
2691 return Err(EngineError::ShapeMismatch(format!(
2692 "kv dims {k_dim}/{v_dim} != n_kv*head_dim {}",
2693 n_kv * head_dim
2694 )));
2695 }
2696 let mut k = self.wmm(wk, &xn, k_dim, hidden, GemmAcct::Attn)?;
2697 let mut v = self.wmm(wv, &xn, v_dim, hidden, GemmAcct::Attn)?;
2698 if let Some(kn) = &attn.k_norm {
2699 if kn.len() != head_dim {
2700 return Err(EngineError::ShapeMismatch(format!(
2701 "k_norm len {} != head_dim {head_dim}",
2702 kn.len()
2703 )));
2704 }
2705 k = self.norm(&k, kn)?;
2706 }
2707 Self::apply_rope(&mut k, head_dim, pos, theta, rope)?;
2708 v = self.apply_v_norm(v, attn.v_norm.as_deref(), head_dim)?;
2709 state.k_caches[li].extend_from_slice(&k);
2710 state.v_caches[li].extend_from_slice(&v);
2711 state.last_kv_src.insert(attn.kind, li);
2712 (li, li)
2713 } else {
2714 let src = state.last_kv_src.get(&attn.kind).copied().ok_or_else(|| {
2715 EngineError::Format(format!(
2716 "KV-consumer layer {li} has no producer of kind {:?}",
2717 attn.kind
2718 ))
2719 })?;
2720 (src, src)
2721 };
2722 let kv_dim = n_kv * head_dim;
2723 let (k_view, v_view) = kv_sliding_view(
2724 &state.k_caches[k_src],
2725 &state.v_caches[v_src],
2726 kv_dim,
2727 self.attn_window(attn.kind),
2728 )?;
2729 let attn_out = attention_with_scale(
2730 &q,
2731 k_view,
2732 v_view,
2733 n_heads,
2734 n_kv,
2735 head_dim,
2736 self.attn_scale(head_dim),
2737 )?;
2738 let mut attn_out = attn_out;
2739 if attn.q_gate {
2740 apply_sigmoid_gate(&mut attn_out, &gate)?;
2741 }
2742 let ao = self.wmm(&attn.wo, &attn_out, hidden, q_dim, GemmAcct::Attn)?;
2743 if let (Some(ref mut streams), Some(preds), Some(laurel)) =
2744 (gemma3n_streams.as_mut(), gemma3n_preds, gemma3n_laurel)
2745 {
2746 self.gemma3n_after_attn(
2747 streams,
2748 &preds,
2749 &laurel,
2750 &ao,
2751 layer,
2752 li,
2753 ple_tok.as_deref(),
2754 hidden,
2755 )?;
2756 continue;
2757 }
2758 self.add_normed_residual(&mut x, &ao, layer.post_attn_norm.as_deref())?;
2759 }
2760 LayerOp::Conv(conv) => {
2761 let bcx = self.wmm(&conv.in_proj, &xn, 3 * hidden, hidden, GemmAcct::Attn)?;
2762 let mut bx = vec![0.0f32; hidden];
2763 let mut c_gate = vec![0.0f32; hidden];
2764 for i in 0..hidden {
2765 let b = bcx[i];
2766 let c = bcx[hidden + i];
2767 let xx = bcx[2 * hidden + i];
2768 bx[i] = b * xx;
2769 c_gate[i] = c;
2770 }
2771 let cstate = state.conv_states[li]
2772 .as_mut()
2773 .ok_or_else(|| EngineError::ShapeMismatch("missing conv state".into()))?;
2774 let conv_y =
2775 short_conv_step(&bx, &conv.kernel, cstate, hidden, conv.kernel_size)?;
2776 let mut y = vec![0.0f32; hidden];
2777 for i in 0..hidden {
2778 y[i] = c_gate[i] * conv_y[i];
2779 }
2780 let ao = self.wmm(&conv.out_proj, &y, hidden, hidden, GemmAcct::Attn)?;
2781 self.add_normed_residual(&mut x, &ao, layer.post_attn_norm.as_deref())?;
2782 }
2783 LayerOp::Linear(dn) => {
2784 let key_dim = dn.n_k_heads * dn.head_k;
2785 let value_dim = dn.n_v_heads * dn.head_v;
2786 let qkvz_out = 2 * key_dim + 2 * value_dim;
2787 let mixed = self.wmm(&dn.qkvz, &xn, qkvz_out, hidden, GemmAcct::Attn)?;
2788 let mut q = mixed[0..key_dim].to_vec();
2789 let mut k = mixed[key_dim..2 * key_dim].to_vec();
2790 let mut v = mixed[2 * key_dim..2 * key_dim + value_dim].to_vec();
2791 let z = mixed[2 * key_dim + value_dim..].to_vec();
2792 let mut qkv = Vec::with_capacity(key_dim * 2 + value_dim);
2793 qkv.extend_from_slice(&q);
2794 qkv.extend_from_slice(&k);
2795 qkv.extend_from_slice(&v);
2796 let conv_dim = qkv.len();
2797 let cstate = state.conv_states[li].as_mut().ok_or_else(|| {
2798 EngineError::ShapeMismatch("missing delta conv state".into())
2799 })?;
2800 let mut mixed_c = short_conv_step(&qkv, &dn.conv, cstate, conv_dim, dn.conv_k)?;
2801 silu_vec(&mut mixed_c);
2802 q.copy_from_slice(&mixed_c[0..key_dim]);
2803 k.copy_from_slice(&mixed_c[key_dim..2 * key_dim]);
2804 v.copy_from_slice(&mixed_c[2 * key_dim..]);
2805 let ba = self.wmm(&dn.ba, &xn, 2 * dn.n_v_heads, hidden, GemmAcct::Attn)?;
2806 let mut beta = vec![0.0f32; dn.n_v_heads];
2807 let mut g = vec![0.0f32; dn.n_v_heads];
2808 for h in 0..dn.n_v_heads {
2809 beta[h] = 1.0 / (1.0 + (-ba[h]).exp());
2810 let alpha =
2811 -dn.a_log[h].exp() * softplus(ba[dn.n_v_heads + h] + dn.dt_bias[h]);
2812 g[h] = alpha.exp();
2813 }
2814 if dn.n_v_heads != dn.n_k_heads {
2815 return Err(EngineError::Unsupported(
2816 "DeltaNet GQA (n_v != n_k) not implemented".into(),
2817 ));
2818 }
2819 let s = state.delta_states[li].as_mut().ok_or_else(|| {
2820 EngineError::ShapeMismatch("missing delta recurrent state".into())
2821 })?;
2822 let mut core = gated_delta_step(GatedDeltaStep {
2823 q: &q,
2824 k: &k,
2825 v: &v,
2826 g: &g,
2827 beta: &beta,
2828 state: s,
2829 n_heads: dn.n_v_heads,
2830 dk: dn.head_k,
2831 dv: dn.head_v,
2832 })?;
2833 core = rms_norm(&core, &dn.out_norm, 1e-6)?;
2835 let mut z_act = z;
2836 silu_vec(&mut z_act);
2837 for i in 0..core.len() {
2838 core[i] *= z_act[i];
2839 }
2840 let ao = self.wmm(&dn.out_proj, &core, hidden, value_dim, GemmAcct::Attn)?;
2841 self.add_normed_residual(&mut x, &ao, layer.post_attn_norm.as_deref())?;
2842 }
2843 }
2844 let xn2 = self.norm(&x, &layer.ffn_norm)?;
2845 let down = self.apply_ffn(layer, &xn2, hidden)?;
2846 self.add_normed_residual(&mut x, &down, layer.post_ffn_norm.as_deref())?;
2847 self.apply_ple(&mut x, layer, li, ple_tok.as_deref(), hidden)?;
2848 Self::apply_layer_scalar(&mut x, layer.layer_scalar);
2849 }
2850 if let Some(streams) = gemma3n_streams {
2851 x = self.gemma3n_unembed_streams(&streams, hidden)?;
2852 }
2853 state.pos += 1;
2854 let xn = self.norm(&x, &self.weights.output_norm)?;
2855 if !self.weights.output.data.len().is_multiple_of(hidden) {
2856 return Err(EngineError::ShapeMismatch(format!(
2857 "lm_head len {} not divisible by hidden {hidden}",
2858 self.weights.output.data.len()
2859 )));
2860 }
2861 let out_rows = self.weights.output.data.len() / hidden;
2862 let logits = self.wmm(
2863 &self.weights.output,
2864 &xn,
2865 out_rows,
2866 hidden,
2867 GemmAcct::LmHead,
2868 )?;
2869 Ok(self.softcap_logits(logits))
2870 }
2871
2872 fn softcap_logits(&self, mut logits: Vec<f32>) -> Vec<f32> {
2873 if let Some(cap) = self.final_logit_softcap.filter(|c| *c > 0.0) {
2874 for x in &mut logits {
2875 *x = (*x / cap).tanh() * cap;
2876 }
2877 }
2878 logits
2879 }
2880}
2881
2882fn gaussian_topk(gate: &[f32], inter: usize, sparsity: f32) -> Result<Vec<f32>, EngineError> {
2884 if inter == 0 || !gate.len().is_multiple_of(inter) {
2885 return Err(EngineError::ShapeMismatch(
2886 "gaussian_topk: gate len not divisible by intermediate_size".into(),
2887 ));
2888 }
2889 let z = if (sparsity - 0.95).abs() < 0.02 {
2890 1.644_853_8
2891 } else {
2892 1.644_853_8 * (sparsity / 0.95).clamp(0.0, 4.0)
2894 };
2895 let seq = gate.len() / inter;
2896 let mut out = vec![0.0f32; gate.len()];
2897 let n = inter as f32;
2898 for t in 0..seq {
2899 let row = &gate[t * inter..(t + 1) * inter];
2900 let mean = row.iter().sum::<f32>() / n;
2901 let mut var = 0.0f32;
2902 for &v in row {
2903 let d = v - mean;
2904 var += d * d;
2905 }
2906 var /= n;
2907 let cutoff = mean + var.sqrt() * z;
2908 for i in 0..inter {
2909 out[t * inter + i] = (row[i] - cutoff).max(0.0);
2910 }
2911 }
2912 Ok(out)
2913}
2914
2915fn argmax(v: &[f32]) -> u32 {
2916 let mut best = 0usize;
2917 let mut best_v = f32::NEG_INFINITY;
2918 for (i, &x) in v.iter().enumerate() {
2919 if x > best_v {
2920 best_v = x;
2921 best = i;
2922 }
2923 }
2924 best as u32
2925}
2926
2927pub fn confidence_from_logits(logits: &[f32]) -> f32 {
2929 if logits.is_empty() {
2930 return 0.0;
2931 }
2932 let m = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
2933 let mut sum = 0.0f32;
2934 let mut maxp = 0.0f32;
2935 for &x in logits {
2936 let e = (x - m).exp();
2937 sum += e;
2938 if e > maxp {
2939 maxp = e;
2940 }
2941 }
2942 if sum > 0.0 {
2943 maxp / sum
2944 } else {
2945 0.0
2946 }
2947}
2948
2949#[allow(dead_code)]
2950pub fn cache_shapes_ok(cache: &HashMap<usize, Vec<f32>>, kv_dim: usize) -> bool {
2951 cache.values().all(|v| v.len().is_multiple_of(kv_dim))
2952}
2953
2954#[cfg(test)]
2955mod tests {
2956 use super::*;
2957 use crate::family::{arch_class_representatives, graph_hook, lookup_family, require_stage_b};
2958 use crate::fixture::write_tiny_q4_bundle;
2959 use aria_kernel::{resolve_compute, ComputePref};
2960 use serde_json::{json, Value};
2961
2962 #[test]
2963 fn gemma4_fills_hub_bundle_missing_geometry_fields() {
2964 let dir = tempfile::tempdir().unwrap();
2965 write_tiny_q4_bundle(dir.path()).unwrap();
2966 let cfg_path = dir.path().join("config.json");
2967 let raw = std::fs::read_to_string(&cfg_path).unwrap();
2968 let mut cfg: Value = serde_json::from_str(&raw).unwrap();
2969 let model = cfg["model"].as_object_mut().unwrap();
2971 for key in [
2972 "layer_types",
2973 "sliding_window",
2974 "partial_rotary_factor",
2975 "global_head_dim",
2976 "head_dim",
2977 ] {
2978 model.remove(key);
2979 }
2980 std::fs::write(&cfg_path, cfg.to_string()).unwrap();
2981 let s = SessionBuilder::new()
2982 .model(dir.path())
2983 .family("gemma/gemma-4-e2b-it")
2984 .build()
2985 .unwrap();
2986 assert_eq!(s.config().sliding_window, Some(512));
2987 assert_eq!(s.config().partial_rotary_factor, Some(0.25));
2988 assert!(s.config().head_dim.unwrap_or(0) > 0);
2989 assert!(s.config().global_head_dim.unwrap_or(0) > 0);
2990 assert_eq!(
2991 s.config().layer_types.as_ref().map(|t| t.len()),
2992 Some(s.config().num_layers)
2993 );
2994 }
2995
2996 #[test]
2997 fn gemma3_text_not_gemma3n() {
2998 assert!(is_gemma3_text("gemma/gemma-3-270m-it"));
2999 assert!(is_gemma3_text("gemma/gemma-3-1b-it"));
3000 assert!(!is_gemma3_text("gemma/gemma-3n-e2b-it"));
3001 assert!(!is_gemma3_text("gemma/gemma-4-e2b-it"));
3002 let t = default_gemma3_layer_types(18);
3003 assert_eq!(t[5], "full_attention");
3004 assert_eq!(t[11], "full_attention");
3005 assert_eq!(t[17], "full_attention");
3006 assert_eq!(t.iter().filter(|s| *s == "sliding_attention").count(), 15);
3007 }
3008
3009 #[test]
3010 fn gemma3_fills_hub_bundle_and_dual_rope() {
3011 let dir = tempfile::tempdir().unwrap();
3012 write_tiny_q4_bundle(dir.path()).unwrap();
3013 let cfg_path = dir.path().join("config.json");
3014 let raw = std::fs::read_to_string(&cfg_path).unwrap();
3015 let mut cfg: Value = serde_json::from_str(&raw).unwrap();
3016 let model = cfg["model"].as_object_mut().unwrap();
3017 for key in ["layer_types", "sliding_window", "hidden_act"] {
3018 model.remove(key);
3019 }
3020 std::fs::write(&cfg_path, cfg.to_string()).unwrap();
3021 let mut s = SessionBuilder::new()
3022 .model(dir.path())
3023 .family("gemma/gemma-3-270m-it")
3024 .build()
3025 .unwrap();
3026 assert_eq!(s.config().sliding_window, Some(512));
3027 assert_eq!(s.config().hidden_act.as_deref(), Some("gelu_pytorch_tanh"));
3028 let types = s.config().layer_types.as_ref().expect("layer_types");
3029 assert_eq!(types.len(), s.config().num_layers);
3030 assert!(types.iter().all(|t| t == "sliding_attention"));
3031 assert_eq!(
3032 s.layer_rope_params(AttnKind::Sliding),
3033 (10_000.0, RopeMode::Full)
3034 );
3035 assert_eq!(
3036 s.layer_rope_params(AttnKind::Full),
3037 (1_000_000.0, RopeMode::Full)
3038 );
3039 let gen = s
3040 .generate(
3041 &[1, 2],
3042 &GenerateOpts {
3043 max_tokens: 2,
3044 temperature: 0.0,
3045 },
3046 )
3047 .unwrap();
3048 assert_eq!(gen.tokens.len(), 2);
3049 }
3050
3051 #[test]
3052 fn gemma3n_not_gemma3_text_and_fills_4plus1_dual_rope() {
3053 assert!(is_gemma3n("gemma/gemma-3n-e2b-it"));
3054 assert!(is_gemma3n("gemma/gemma-3n-e4b-it"));
3055 assert!(!is_gemma3n("gemma/gemma-3-270m-it"));
3056 let t = default_gemma4_layer_types(30);
3057 assert_eq!(t[4], "full_attention");
3058 assert_eq!(t[9], "full_attention");
3059 assert_eq!(t[29], "full_attention");
3060 assert_eq!(t.iter().filter(|s| *s == "sliding_attention").count(), 24);
3061 assert_eq!(gemma3n_default_kv_shared(30), 10);
3062 assert_eq!(gemma3n_default_kv_shared(35), 15);
3063 assert_eq!(gemma3n_default_kv_shared(2), 0);
3064 assert!(
3065 (altup_router_input_scale(2048) - 1.0 / 2048.0).abs() < 1e-12,
3066 "HF router_input_scale is 1/hidden, not 1/sqrt(hidden)"
3067 );
3068
3069 let dir = tempfile::tempdir().unwrap();
3070 write_tiny_q4_bundle(dir.path()).unwrap();
3071 let cfg_path = dir.path().join("config.json");
3072 let raw = std::fs::read_to_string(&cfg_path).unwrap();
3073 let mut cfg: Value = serde_json::from_str(&raw).unwrap();
3074 let model = cfg["model"].as_object_mut().unwrap();
3075 for key in ["layer_types", "sliding_window", "hidden_act"] {
3076 model.remove(key);
3077 }
3078 std::fs::write(&cfg_path, cfg.to_string()).unwrap();
3079 let mut s = SessionBuilder::new()
3080 .model(dir.path())
3081 .family("gemma/gemma-3n-e2b-it")
3082 .build()
3083 .unwrap();
3084 assert!(!s.use_gemma_norm, "Gemma-3n RMSNorm is *w, not *(1+w)");
3085 assert!((s.attn_scale(256) - 1.0).abs() < 1e-6);
3086 assert_eq!(s.final_logit_softcap, Some(30.0));
3087 assert_eq!(s.config().sliding_window, Some(512));
3088 assert_eq!(
3089 s.layer_rope_params(AttnKind::Sliding),
3090 (10_000.0, RopeMode::Full)
3091 );
3092 assert_eq!(
3093 s.layer_rope_params(AttnKind::Full),
3094 (1_000_000.0, RopeMode::Full)
3095 );
3096 let gen = s
3097 .generate(
3098 &[1, 2],
3099 &GenerateOpts {
3100 max_tokens: 2,
3101 temperature: 0.0,
3102 },
3103 )
3104 .unwrap();
3105 assert_eq!(gen.tokens.len(), 2);
3106 }
3107
3108 #[test]
3109 fn generate_tokens() {
3110 let dir = tempfile::tempdir().unwrap();
3111 write_tiny_q4_bundle(dir.path()).unwrap();
3112 let mut s = SessionBuilder::new()
3113 .model(dir.path())
3114 .family("gemma/gemma-4-e2b-it")
3115 .build()
3116 .unwrap();
3117 assert_eq!(s.config().sliding_window, Some(512));
3118 assert_eq!(s.config().partial_rotary_factor, Some(0.25));
3119 assert_eq!(s.config().head_dim, Some(16));
3120 assert_eq!(s.config().global_head_dim, Some(16));
3121 assert_eq!(
3122 s.config().layer_types,
3123 Some(vec!["full_attention".into(), "full_attention".into()])
3124 );
3125 assert_eq!(
3126 s.layer_rope_params(AttnKind::Full),
3127 (1_000_000.0, RopeMode::Proportional(0.25))
3128 );
3129 assert_eq!(
3130 s.layer_rope_params(AttnKind::Sliding),
3131 (10_000.0, RopeMode::Full)
3132 );
3133 assert_eq!(s.attn_window(AttnKind::Sliding), Some(512));
3134 assert_eq!(s.attn_window(AttnKind::Full), None);
3135 let prompt = s.encode_text("hi");
3136 let gen = s
3137 .generate(
3138 &prompt,
3139 &GenerateOpts {
3140 max_tokens: 4,
3141 temperature: 0.0,
3142 },
3143 )
3144 .unwrap();
3145 assert!(!gen.tokens.is_empty());
3146 assert!(!gen.text.is_empty());
3147 }
3148
3149 #[test]
3150 fn split_interleaved_q_gate_matches_hf_chunk() {
3151 let mixed = vec![1.0, 2.0, 10.0, 20.0, 3.0, 4.0, 30.0, 40.0];
3153 let (q, g) = split_interleaved_q_gate(&mixed, 1, 2, 2).unwrap();
3154 assert_eq!(q, vec![1.0, 2.0, 3.0, 4.0]);
3155 assert_eq!(g, vec![10.0, 20.0, 30.0, 40.0]);
3156 }
3157
3158 #[test]
3159 fn qwen35_attn_output_gate_generate() {
3160 let dir = tempfile::tempdir().unwrap();
3161 let hidden = 8usize;
3162 let inter = 16usize;
3163 let vocab = 16usize;
3164 let n_heads = 2usize;
3165 let n_kv = 1usize;
3166 let head_dim = 4usize;
3167 let q_dim = n_heads * head_dim;
3168 let k_dim = n_kv * head_dim;
3169
3170 let mut tensors = serde_json::Map::new();
3171 let mut bin = Vec::new();
3172 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
3173 let offset = bin.len();
3174 for &v in data {
3175 bin.extend_from_slice(&v.to_le_bytes());
3176 }
3177 let nbytes = data.len() * 4;
3178 let mut meta = serde_json::Map::new();
3179 meta.insert("kind".into(), json!("raw"));
3180 meta.insert("dtype".into(), json!("f32"));
3181 meta.insert("shape".into(), json!(shape));
3182 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
3183 tensors.insert(name.to_string(), Value::Object(meta));
3184 };
3185 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
3186 add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
3187 let n1 = vec![1.0f32; hidden];
3188 add_raw("model.layers.0.input_layernorm.weight", vec![hidden], &n1);
3189 add_raw(
3190 "model.layers.0.post_attention_layernorm.weight",
3191 vec![hidden],
3192 &n1,
3193 );
3194 let wq = vec![0.02f32; 2 * q_dim * hidden];
3195 let wk = vec![0.02f32; k_dim * hidden];
3196 let wv = vec![0.02f32; k_dim * hidden];
3197 let wo = vec![0.02f32; hidden * q_dim];
3198 add_raw(
3199 "model.layers.0.self_attn.q_proj.weight",
3200 vec![2 * q_dim, hidden],
3201 &wq,
3202 );
3203 add_raw(
3204 "model.layers.0.self_attn.k_proj.weight",
3205 vec![k_dim, hidden],
3206 &wk,
3207 );
3208 add_raw(
3209 "model.layers.0.self_attn.v_proj.weight",
3210 vec![k_dim, hidden],
3211 &wv,
3212 );
3213 add_raw(
3214 "model.layers.0.self_attn.o_proj.weight",
3215 vec![hidden, q_dim],
3216 &wo,
3217 );
3218 let qn = vec![1.0f32; head_dim];
3219 add_raw(
3220 "model.layers.0.self_attn.q_norm.weight",
3221 vec![head_dim],
3222 &qn,
3223 );
3224 add_raw(
3225 "model.layers.0.self_attn.k_norm.weight",
3226 vec![head_dim],
3227 &qn,
3228 );
3229 let g = vec![0.02f32; inter * hidden];
3230 let d = vec![0.02f32; hidden * inter];
3231 add_raw(
3232 "model.layers.0.mlp.gate_proj.weight",
3233 vec![inter, hidden],
3234 &g,
3235 );
3236 add_raw("model.layers.0.mlp.up_proj.weight", vec![inter, hidden], &g);
3237 add_raw(
3238 "model.layers.0.mlp.down_proj.weight",
3239 vec![hidden, inter],
3240 &d,
3241 );
3242 add_raw("model.norm.weight", vec![hidden], &n1);
3243 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
3244 let cfg = json!({
3245 "format": "aria-quant-bundle",
3246 "format_version": 2,
3247 "quantization": "test",
3248 "hadamard_seed": 0,
3249 "model": {
3250 "hidden_size": hidden,
3251 "num_layers": 1,
3252 "num_attention_heads": n_heads,
3253 "num_kv_heads": n_kv,
3254 "head_dim": head_dim,
3255 "intermediate_size": inter,
3256 "vocab_size": vocab,
3257 "context_length": 32,
3258 "rope_theta": 10000.0,
3259 "layer_types": ["full_attention"]
3260 },
3261 "tensors": tensors
3262 });
3263 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
3264 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
3265 let mut s = SessionBuilder::new()
3266 .model(dir.path())
3267 .family("qwen/qwen3-0.6b")
3268 .build()
3269 .unwrap();
3270 let gen = s
3271 .generate(
3272 &[1, 2],
3273 &GenerateOpts {
3274 max_tokens: 2,
3275 temperature: 0.0,
3276 },
3277 )
3278 .unwrap();
3279 assert_eq!(gen.tokens.len(), 2);
3280 }
3281
3282 #[test]
3283 fn materialize_accepts_hf_tensor_names() {
3284 let dir = tempfile::tempdir().unwrap();
3286 let hidden = 8usize;
3287 let layers = 1usize;
3288 let inter = 16usize;
3289 let vocab = 16usize;
3290 let n_heads = 2usize;
3291 let n_kv = 1usize;
3292 let head_dim = 4usize; let q_dim = n_heads * head_dim;
3294 let k_dim = n_kv * head_dim;
3295
3296 let mut tensors = serde_json::Map::new();
3297 let mut bin = Vec::new();
3298 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
3299 let offset = bin.len();
3300 for &v in data {
3301 bin.extend_from_slice(&v.to_le_bytes());
3302 }
3303 let nbytes = data.len() * 4;
3304 let mut meta = serde_json::Map::new();
3305 meta.insert("kind".into(), json!("raw"));
3306 meta.insert("dtype".into(), json!("f32"));
3307 meta.insert("shape".into(), json!(shape));
3308 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
3309 tensors.insert(name.to_string(), Value::Object(meta));
3310 };
3311 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
3312 add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
3313 let n1 = vec![1.0f32; hidden];
3314 add_raw("model.layers.0.input_layernorm.weight", vec![hidden], &n1);
3315 add_raw(
3316 "model.layers.0.post_attention_layernorm.weight",
3317 vec![hidden],
3318 &n1,
3319 );
3320 let wq = vec![0.01f32; q_dim * hidden];
3321 let wk = vec![0.01f32; k_dim * hidden];
3322 let wv = vec![0.01f32; k_dim * hidden];
3323 let wo = vec![0.01f32; hidden * q_dim];
3324 add_raw(
3325 "model.layers.0.self_attn.q_proj.weight",
3326 vec![q_dim, hidden],
3327 &wq,
3328 );
3329 add_raw(
3330 "model.layers.0.self_attn.k_proj.weight",
3331 vec![k_dim, hidden],
3332 &wk,
3333 );
3334 add_raw(
3335 "model.layers.0.self_attn.v_proj.weight",
3336 vec![k_dim, hidden],
3337 &wv,
3338 );
3339 add_raw(
3340 "model.layers.0.self_attn.o_proj.weight",
3341 vec![hidden, q_dim],
3342 &wo,
3343 );
3344 let g = vec![0.01f32; inter * hidden];
3345 let d = vec![0.01f32; hidden * inter];
3346 add_raw(
3347 "model.layers.0.mlp.gate_proj.weight",
3348 vec![inter, hidden],
3349 &g,
3350 );
3351 add_raw("model.layers.0.mlp.up_proj.weight", vec![inter, hidden], &g);
3352 add_raw(
3353 "model.layers.0.mlp.down_proj.weight",
3354 vec![hidden, inter],
3355 &d,
3356 );
3357 add_raw("model.norm.weight", vec![hidden], &n1);
3358 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
3359
3360 let cfg = json!({
3361 "format": "aria-quant-bundle",
3362 "format_version": 2,
3363 "quantization": "test",
3364 "group_size_default": 32,
3365 "hadamard_seed": 0,
3366 "model": {
3367 "hidden_size": hidden,
3368 "num_layers": layers,
3369 "num_attention_heads": n_heads,
3370 "num_kv_heads": n_kv,
3371 "intermediate_size": inter,
3372 "vocab_size": vocab,
3373 "context_length": 32,
3374 "rope_theta": 10000.0
3375 },
3376 "tensors": tensors
3377 });
3378 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
3379 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
3380
3381 let mut s = SessionBuilder::new()
3382 .model(dir.path())
3383 .family("qwen/qwen3-0.6b")
3384 .build()
3385 .unwrap();
3386 let gen = s
3387 .generate(
3388 &[1, 2],
3389 &GenerateOpts {
3390 max_tokens: 2,
3391 temperature: 0.0,
3392 },
3393 )
3394 .unwrap();
3395 assert_eq!(gen.tokens.len(), 2);
3396 }
3397
3398 #[test]
3399 fn materialize_accepts_language_model_prefix_and_pre_ffn_norm() {
3400 let dir = tempfile::tempdir().unwrap();
3402 let hidden = 8usize;
3403 let layers = 1usize;
3404 let inter = 16usize;
3405 let vocab = 16usize;
3406 let n_heads = 2usize;
3407 let n_kv = 1usize;
3408 let head_dim = 4usize;
3409 let q_dim = n_heads * head_dim;
3410 let k_dim = n_kv * head_dim;
3411
3412 let mut tensors = serde_json::Map::new();
3413 let mut bin = Vec::new();
3414 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
3415 let offset = bin.len();
3416 for &v in data {
3417 bin.extend_from_slice(&v.to_le_bytes());
3418 }
3419 let nbytes = data.len() * 4;
3420 let mut meta = serde_json::Map::new();
3421 meta.insert("kind".into(), json!("raw"));
3422 meta.insert("dtype".into(), json!("f32"));
3423 meta.insert("shape".into(), json!(shape));
3424 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
3425 tensors.insert(name.to_string(), Value::Object(meta));
3426 };
3427 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
3428 let p = "model.language_model";
3429 add_raw(
3430 &format!("{p}.embed_tokens.weight"),
3431 vec![vocab, hidden],
3432 &emb,
3433 );
3434 let n1 = vec![1.0f32; hidden];
3435 add_raw(
3436 &format!("{p}.layers.0.input_layernorm.weight"),
3437 vec![hidden],
3438 &n1,
3439 );
3440 add_raw(
3441 &format!("{p}.layers.0.pre_feedforward_layernorm.weight"),
3442 vec![hidden],
3443 &n1,
3444 );
3445 let wq = vec![0.01f32; q_dim * hidden];
3446 let wk = vec![0.01f32; k_dim * hidden];
3447 let wv = vec![0.01f32; k_dim * hidden];
3448 let wo = vec![0.01f32; hidden * q_dim];
3449 add_raw(
3450 &format!("{p}.layers.0.self_attn.q_proj.weight"),
3451 vec![q_dim, hidden],
3452 &wq,
3453 );
3454 add_raw(
3455 &format!("{p}.layers.0.self_attn.k_proj.weight"),
3456 vec![k_dim, hidden],
3457 &wk,
3458 );
3459 add_raw(
3460 &format!("{p}.layers.0.self_attn.v_proj.weight"),
3461 vec![k_dim, hidden],
3462 &wv,
3463 );
3464 add_raw(
3465 &format!("{p}.layers.0.self_attn.o_proj.weight"),
3466 vec![hidden, q_dim],
3467 &wo,
3468 );
3469 let g = vec![0.01f32; inter * hidden];
3470 let d = vec![0.01f32; hidden * inter];
3471 add_raw(
3472 &format!("{p}.layers.0.mlp.gate_proj.weight"),
3473 vec![inter, hidden],
3474 &g,
3475 );
3476 add_raw(
3477 &format!("{p}.layers.0.mlp.up_proj.weight"),
3478 vec![inter, hidden],
3479 &g,
3480 );
3481 add_raw(
3482 &format!("{p}.layers.0.mlp.down_proj.weight"),
3483 vec![hidden, inter],
3484 &d,
3485 );
3486 add_raw(&format!("{p}.norm.weight"), vec![hidden], &n1);
3487 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
3488
3489 let cfg = json!({
3490 "format": "aria-quant-bundle",
3491 "format_version": 2,
3492 "quantization": "test",
3493 "group_size_default": 32,
3494 "hadamard_seed": 0,
3495 "model": {
3496 "hidden_size": hidden,
3497 "num_layers": layers,
3498 "num_attention_heads": n_heads,
3499 "num_kv_heads": n_kv,
3500 "intermediate_size": inter,
3501 "vocab_size": vocab,
3502 "context_length": 32,
3503 "rope_theta": 10000.0,
3504 "head_dim": head_dim,
3505 "global_head_dim": head_dim,
3506 "sliding_window": 512,
3507 "partial_rotary_factor": 0.25,
3508 "layer_types": ["full_attention"]
3509 },
3510 "tensors": tensors
3511 });
3512 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
3513 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
3514
3515 let mut s = SessionBuilder::new()
3516 .model(dir.path())
3517 .family("gemma/gemma-4-e2b-it")
3518 .build()
3519 .unwrap();
3520 let gen = s
3521 .generate(
3522 &[1, 2],
3523 &GenerateOpts {
3524 max_tokens: 2,
3525 temperature: 0.0,
3526 },
3527 )
3528 .unwrap();
3529 assert_eq!(gen.tokens.len(), 2);
3530 }
3531
3532 #[test]
3533 fn gemma4_style_double_wide_mlp_and_shared_kv() {
3534 let dir = tempfile::tempdir().unwrap();
3537 let hidden = 8usize;
3538 let layers = 2usize;
3539 let inter = 16usize;
3540 let inter_wide = 32usize;
3541 let vocab = 16usize;
3542 let n_heads = 2usize;
3543 let n_kv = 1usize;
3544 let head_dim = 4usize;
3545 let q_dim = n_heads * head_dim;
3546 let k_dim = n_kv * head_dim;
3547 let p = "model.language_model";
3548
3549 let mut tensors = serde_json::Map::new();
3550 let mut bin = Vec::new();
3551 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
3552 let offset = bin.len();
3553 for &v in data {
3554 bin.extend_from_slice(&v.to_le_bytes());
3555 }
3556 let nbytes = data.len() * 4;
3557 let mut meta = serde_json::Map::new();
3558 meta.insert("kind".into(), json!("raw"));
3559 meta.insert("dtype".into(), json!("f32"));
3560 meta.insert("shape".into(), json!(shape));
3561 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
3562 tensors.insert(name.to_string(), Value::Object(meta));
3563 };
3564 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
3565 add_raw(
3566 &format!("{p}.embed_tokens.weight"),
3567 vec![vocab, hidden],
3568 &emb,
3569 );
3570 let n1 = vec![1.0f32; hidden];
3571 let wq = vec![0.01f32; q_dim * hidden];
3572 let wk = vec![0.01f32; k_dim * hidden];
3573 let wv = vec![0.01f32; k_dim * hidden];
3574 let wo = vec![0.01f32; hidden * q_dim];
3575 for li in 0..layers {
3576 let layer_inter = if li == 0 { inter } else { inter_wide };
3577 add_raw(
3578 &format!("{p}.layers.{li}.input_layernorm.weight"),
3579 vec![hidden],
3580 &n1,
3581 );
3582 add_raw(
3583 &format!("{p}.layers.{li}.pre_feedforward_layernorm.weight"),
3584 vec![hidden],
3585 &n1,
3586 );
3587 add_raw(
3588 &format!("{p}.layers.{li}.self_attn.q_proj.weight"),
3589 vec![q_dim, hidden],
3590 &wq,
3591 );
3592 if li == 0 {
3593 add_raw(
3594 &format!("{p}.layers.{li}.self_attn.k_proj.weight"),
3595 vec![k_dim, hidden],
3596 &wk,
3597 );
3598 add_raw(
3599 &format!("{p}.layers.{li}.self_attn.v_proj.weight"),
3600 vec![k_dim, hidden],
3601 &wv,
3602 );
3603 }
3604 add_raw(
3605 &format!("{p}.layers.{li}.self_attn.o_proj.weight"),
3606 vec![hidden, q_dim],
3607 &wo,
3608 );
3609 let g = vec![0.01f32; layer_inter * hidden];
3610 let d = vec![0.01f32; hidden * layer_inter];
3611 add_raw(
3612 &format!("{p}.layers.{li}.mlp.gate_proj.weight"),
3613 vec![layer_inter, hidden],
3614 &g,
3615 );
3616 add_raw(
3617 &format!("{p}.layers.{li}.mlp.up_proj.weight"),
3618 vec![layer_inter, hidden],
3619 &g,
3620 );
3621 add_raw(
3622 &format!("{p}.layers.{li}.mlp.down_proj.weight"),
3623 vec![hidden, layer_inter],
3624 &d,
3625 );
3626 }
3627 add_raw(&format!("{p}.norm.weight"), vec![hidden], &n1);
3628 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
3629
3630 let cfg = json!({
3631 "format": "aria-quant-bundle",
3632 "format_version": 2,
3633 "quantization": "test",
3634 "group_size_default": 32,
3635 "hadamard_seed": 0,
3636 "model": {
3637 "hidden_size": hidden,
3638 "num_layers": layers,
3639 "num_attention_heads": n_heads,
3640 "num_kv_heads": n_kv,
3641 "intermediate_size": inter,
3642 "vocab_size": vocab,
3643 "context_length": 32,
3644 "rope_theta": 10000.0,
3645 "num_kv_shared_layers": 1,
3646 "head_dim": head_dim,
3647 "global_head_dim": head_dim,
3648 "sliding_window": 512,
3649 "partial_rotary_factor": 0.25,
3650 "layer_types": ["full_attention", "full_attention"]
3651 },
3652 "tensors": tensors
3653 });
3654 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
3655 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
3656
3657 let mut s = SessionBuilder::new()
3658 .model(dir.path())
3659 .family("gemma/gemma-4-e2b-it")
3660 .build()
3661 .unwrap();
3662 let gen = s
3663 .generate(
3664 &[1, 2],
3665 &GenerateOpts {
3666 max_tokens: 2,
3667 temperature: 0.0,
3668 },
3669 )
3670 .unwrap();
3671 assert_eq!(gen.tokens.len(), 2);
3672 }
3673
3674 #[test]
3675 fn stage_b_arch_classes_generate() {
3676 for (path, arch) in arch_class_representatives() {
3677 if matches!(arch, ArchClass::VL | ArchClass::VLA | ArchClass::TextMoE) {
3678 continue; }
3680 if path.contains("qwen3.5") || path.contains("bonsai") {
3682 let dir = tempfile::tempdir().unwrap();
3683 write_tiny_q4_bundle(dir.path()).unwrap();
3684 let err = SessionBuilder::new()
3685 .model(dir.path())
3686 .family(*path)
3687 .build()
3688 .unwrap_err();
3689 assert!(
3690 matches!(err, EngineError::Unsupported(_)),
3691 "{path}: {err:?}"
3692 );
3693 continue;
3694 }
3695 assert!(require_stage_b(path).is_ok(), "{path}");
3696 let dir = tempfile::tempdir().unwrap();
3697 write_tiny_q4_bundle(dir.path()).unwrap();
3698 let mut s = SessionBuilder::new()
3699 .model(dir.path())
3700 .family(*path)
3701 .build()
3702 .unwrap();
3703 assert_eq!(s.arch(), *arch);
3704 assert!(!s.graph_hook_name().is_empty());
3705 let gen = s
3706 .generate(
3707 &s.encode_text("ok"),
3708 &GenerateOpts {
3709 max_tokens: 2,
3710 temperature: 0.0,
3711 },
3712 )
3713 .unwrap();
3714 assert!(!gen.tokens.is_empty(), "{path}");
3715 }
3716 }
3717
3718 #[test]
3719 fn stage_c_vl_vla_hooks() {
3720 let dir = tempfile::tempdir().unwrap();
3721 write_tiny_q4_bundle(dir.path()).unwrap();
3722 let s = SessionBuilder::new()
3723 .model(dir.path())
3724 .family("lfm/lfm2-vl-450m")
3725 .build()
3726 .unwrap();
3727 let rgb = vec![10u8; 3 * 4 * 4];
3728 let err = s.vision_prefix(&rgb, 4, 4).unwrap_err();
3729 assert!(matches!(err, EngineError::Unsupported(_)));
3730
3731 let vla = SessionBuilder::new()
3732 .model(dir.path())
3733 .family("openvla/openvla-7b")
3734 .build()
3735 .unwrap();
3736 let err = vla.predict_action("move", 7).unwrap_err();
3737 assert!(matches!(err, EngineError::Unsupported(_)));
3738 let emb = vla.embed_text("hello").unwrap();
3739 assert_eq!(emb.len(), vla.config().hidden_size);
3740 }
3741
3742 #[test]
3743 fn unknown_family() {
3744 let err = SessionBuilder::new()
3745 .model("/tmp")
3746 .family("no/such-model")
3747 .build()
3748 .unwrap_err();
3749 assert!(matches!(err, EngineError::UnsupportedFamily(_)));
3750 }
3751
3752 #[test]
3753 fn greedy_deterministic() {
3754 let dir = tempfile::tempdir().unwrap();
3755 write_tiny_q4_bundle(dir.path()).unwrap();
3756 let mut s = SessionBuilder::new()
3757 .model(dir.path())
3758 .family("gemma/gemma-4-e2b-it")
3759 .build()
3760 .unwrap();
3761 let prompt = s.encode_text("hi");
3762 let opts = GenerateOpts {
3763 max_tokens: 3,
3764 temperature: 0.0,
3765 };
3766 let a = s.generate(&prompt, &opts).unwrap();
3767 let b = s.generate(&prompt, &opts).unwrap();
3768 assert_eq!(a.tokens, b.tokens);
3769 assert_eq!(a.tokens.len(), 3);
3770 }
3771
3772 #[test]
3773 fn encode_chat_is_longer_than_raw_user_text() {
3774 let dir = tempfile::tempdir().unwrap();
3775 write_tiny_q4_bundle(dir.path()).unwrap();
3776 let s = SessionBuilder::new()
3777 .model(dir.path())
3778 .family("qwen/qwen3-0.6b")
3779 .build()
3780 .unwrap();
3781 let raw = s.encode_text("Hello");
3782 let chat = s.encode_chat(&[ChatTurn::new("user", "Hello")]);
3783 assert!(
3784 chat.len() > raw.len(),
3785 "chat template should wrap the user turn (raw={}, chat={})",
3786 raw.len(),
3787 chat.len()
3788 );
3789 assert!(
3790 (s.config().rope_theta - 1_000_000.0).abs() < 1.0,
3791 "Qwen3 must not keep Llama-default rope_theta=10000, got {}",
3792 s.config().rope_theta
3793 );
3794 }
3795
3796 #[test]
3797 fn incremental_decode_matches_full_recompute() {
3798 let dir = tempfile::tempdir().unwrap();
3799 write_tiny_q4_bundle(dir.path()).unwrap();
3800 let mut s = SessionBuilder::new()
3801 .model(dir.path())
3802 .family("gemma/gemma-4-e2b-it")
3803 .build()
3804 .unwrap();
3805 let prompt = s.encode_text("hi");
3806 let max_tokens = 5usize;
3807
3808 let mut prefix = prompt.clone();
3810 if prefix.is_empty() {
3811 prefix.push(1);
3812 }
3813 let mut full_tokens = Vec::new();
3814 for _ in 0..max_tokens {
3815 let logits = s.forward(&prefix).unwrap();
3816 let next = argmax(&logits);
3817 full_tokens.push(next);
3818 prefix.push(next);
3819 if s.is_stop_id(next) {
3820 full_tokens.pop();
3821 break;
3822 }
3823 }
3824
3825 let incr = s
3826 .generate(
3827 &prompt,
3828 &GenerateOpts {
3829 max_tokens,
3830 temperature: 0.0,
3831 },
3832 )
3833 .unwrap();
3834 assert_eq!(
3835 incr.tokens, full_tokens,
3836 "incremental decode must match full-recompute greedy tokens"
3837 );
3838 }
3839
3840 #[test]
3841 fn profile_records_load_and_generate() {
3842 let dir = tempfile::tempdir().unwrap();
3843 write_tiny_q4_bundle(dir.path()).unwrap();
3844 let mut s = SessionBuilder::new()
3845 .model(dir.path())
3846 .family("gemma/gemma-4-e2b-it")
3847 .compute(ComputePref::Cpu)
3848 .profile(true)
3849 .build()
3850 .unwrap();
3851 assert!(s.compute_label().contains("cpu"));
3852 let load = s.last_profile().expect("load profile");
3853 assert!(!load.ci_fail);
3854 assert!(load.load.materialize_ms >= 0.0);
3855 s.generate(
3856 &s.encode_text("hi"),
3857 &GenerateOpts {
3858 max_tokens: 2,
3859 temperature: 0.0,
3860 },
3861 )
3862 .unwrap();
3863 let p = s.last_profile().expect("generate profile");
3864 let g = p.generate.as_ref().expect("generate timings");
3865 assert!(g.prefill_ms >= 0.0);
3866 assert!(g.decode_ms >= 0.0);
3867 }
3868
3869 #[test]
3870 fn cuda_greedy_matches_cpu_if_available() {
3871 if resolve_compute(ComputePref::Cuda).is_err() {
3872 return;
3873 }
3874 let dir = tempfile::tempdir().unwrap();
3875 write_tiny_q4_bundle(dir.path()).unwrap();
3876 let prompt_text = "hi";
3877 let opts = GenerateOpts {
3878 max_tokens: 4,
3879 temperature: 0.0,
3880 };
3881 let mut cpu = SessionBuilder::new()
3882 .model(dir.path())
3883 .family("gemma/gemma-4-e2b-it")
3884 .compute(ComputePref::Cpu)
3885 .build()
3886 .unwrap();
3887 let mut gpu = SessionBuilder::new()
3888 .model(dir.path())
3889 .family("gemma/gemma-4-e2b-it")
3890 .compute(ComputePref::Cuda)
3891 .build()
3892 .unwrap();
3893 assert!(gpu.compute_label().contains("cuda"));
3894 let prompt = cpu.encode_text(prompt_text);
3895 let a = cpu.generate(&prompt, &opts).unwrap();
3896 let b = gpu.generate(&prompt, &opts).unwrap();
3897 assert_eq!(
3898 a.tokens, b.tokens,
3899 "CUDA greedy tokens must match CPU (tiny bundle)"
3900 );
3901 }
3902
3903 #[test]
3904 fn max_tokens_zero_rejected() {
3905 let dir = tempfile::tempdir().unwrap();
3906 write_tiny_q4_bundle(dir.path()).unwrap();
3907 let mut s = SessionBuilder::new()
3908 .model(dir.path())
3909 .family("gemma/gemma-4-e2b-it")
3910 .build()
3911 .unwrap();
3912 let err = s
3913 .generate(
3914 &s.encode_text("x"),
3915 &GenerateOpts {
3916 max_tokens: 0,
3917 temperature: 0.0,
3918 },
3919 )
3920 .unwrap_err();
3921 assert!(matches!(err, EngineError::InvalidParam(_)));
3922 }
3923
3924 #[test]
3925 fn moe_family_refuses_dense_stub() {
3926 let dir = tempfile::tempdir().unwrap();
3927 write_tiny_q4_bundle(dir.path()).unwrap();
3928 let err = SessionBuilder::new()
3929 .model(dir.path())
3930 .family("lfm/lfm2-8b-a1b")
3931 .build()
3932 .unwrap_err();
3933 assert!(matches!(err, EngineError::Unsupported(_)));
3934 assert_eq!(
3935 lookup_family("lfm/lfm2-8b-a1b").unwrap().arch,
3936 ArchClass::TextMoE
3937 );
3938 assert_eq!(graph_hook(ArchClass::TextMoE), "text_moe_decoder");
3939 }
3940
3941 #[test]
3942 fn geometry_gates_conv_and_experts() {
3943 let dir = tempfile::tempdir().unwrap();
3945 write_tiny_q4_bundle(dir.path()).unwrap();
3946 let cfg_path = dir.path().join("config.json");
3947 let mut cfg: Value =
3948 serde_json::from_str(&std::fs::read_to_string(&cfg_path).unwrap()).unwrap();
3949 cfg["model"]["layer_types"] = json!(["conv", "full_attention"]);
3950 std::fs::write(&cfg_path, cfg.to_string()).unwrap();
3951 let err = SessionBuilder::new()
3952 .model(dir.path())
3953 .family("lfm/lfm2-350m")
3954 .build()
3955 .unwrap_err();
3956 assert!(
3957 matches!(err, EngineError::Format(_)),
3958 "expected missing conv tensors, got {err:?}"
3959 );
3960
3961 let dir2 = tempfile::tempdir().unwrap();
3963 write_tiny_q4_bundle(dir2.path()).unwrap();
3964 let cfg_path = dir2.path().join("config.json");
3965 let mut cfg: Value =
3966 serde_json::from_str(&std::fs::read_to_string(&cfg_path).unwrap()).unwrap();
3967 cfg["model"]["layer_types"] = json!(["linear_attention", "full_attention"]);
3968 std::fs::write(&cfg_path, cfg.to_string()).unwrap();
3969 let err = SessionBuilder::new()
3970 .model(dir2.path())
3971 .family("gemma/gemma-3-270m-it")
3972 .build()
3973 .unwrap_err();
3974 assert!(
3975 matches!(err, EngineError::Format(_)),
3976 "expected missing DeltaNet tensors, got {err:?}"
3977 );
3978 }
3979
3980 #[test]
3981 fn lfm_short_conv_and_attn_generate() {
3982 let dir = tempfile::tempdir().unwrap();
3983 let hidden = 8usize;
3984 let layers = 2usize;
3985 let inter = 16usize;
3986 let vocab = 16usize;
3987 let n_heads = 2usize;
3988 let n_kv = 1usize;
3989 let head_dim = 4usize;
3990 let q_dim = n_heads * head_dim;
3991 let k_dim = n_kv * head_dim;
3992 let kernel = 3usize;
3993
3994 let mut tensors = serde_json::Map::new();
3995 let mut bin = Vec::new();
3996 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
3997 let offset = bin.len();
3998 for &v in data {
3999 bin.extend_from_slice(&v.to_le_bytes());
4000 }
4001 let nbytes = data.len() * 4;
4002 let mut meta = serde_json::Map::new();
4003 meta.insert("kind".into(), json!("raw"));
4004 meta.insert("dtype".into(), json!("f32"));
4005 meta.insert("shape".into(), json!(shape));
4006 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
4007 tensors.insert(name.to_string(), Value::Object(meta));
4008 };
4009 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
4010 add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
4011 let n1 = vec![1.0f32; hidden];
4012 add_raw("model.layers.0.operator_norm.weight", vec![hidden], &n1);
4014 add_raw("model.layers.0.ffn_norm.weight", vec![hidden], &n1);
4015 let in_proj = vec![0.02f32; 3 * hidden * hidden];
4016 let out_proj = vec![0.02f32; hidden * hidden];
4017 let conv_w = vec![0.1f32; hidden * kernel];
4018 add_raw(
4019 "model.layers.0.conv.in_proj.weight",
4020 vec![3 * hidden, hidden],
4021 &in_proj,
4022 );
4023 add_raw(
4024 "model.layers.0.conv.out_proj.weight",
4025 vec![hidden, hidden],
4026 &out_proj,
4027 );
4028 add_raw(
4029 "model.layers.0.conv.conv.weight",
4030 vec![hidden, kernel],
4031 &conv_w,
4032 );
4033 let g = vec![0.02f32; inter * hidden];
4034 let d = vec![0.02f32; hidden * inter];
4035 add_raw(
4036 "model.layers.0.mlp.gate_proj.weight",
4037 vec![inter, hidden],
4038 &g,
4039 );
4040 add_raw("model.layers.0.mlp.up_proj.weight", vec![inter, hidden], &g);
4041 add_raw(
4042 "model.layers.0.mlp.down_proj.weight",
4043 vec![hidden, inter],
4044 &d,
4045 );
4046 add_raw("model.layers.1.operator_norm.weight", vec![hidden], &n1);
4048 add_raw(
4049 "model.layers.1.post_attention_layernorm.weight",
4050 vec![hidden],
4051 &n1,
4052 );
4053 let wq = vec![0.02f32; q_dim * hidden];
4054 let wk = vec![0.02f32; k_dim * hidden];
4055 let wv = vec![0.02f32; k_dim * hidden];
4056 let wo = vec![0.02f32; hidden * q_dim];
4057 add_raw(
4058 "model.layers.1.self_attn.q_proj.weight",
4059 vec![q_dim, hidden],
4060 &wq,
4061 );
4062 add_raw(
4063 "model.layers.1.self_attn.k_proj.weight",
4064 vec![k_dim, hidden],
4065 &wk,
4066 );
4067 add_raw(
4068 "model.layers.1.self_attn.v_proj.weight",
4069 vec![k_dim, hidden],
4070 &wv,
4071 );
4072 add_raw(
4073 "model.layers.1.self_attn.o_proj.weight",
4074 vec![hidden, q_dim],
4075 &wo,
4076 );
4077 add_raw(
4078 "model.layers.1.mlp.gate_proj.weight",
4079 vec![inter, hidden],
4080 &g,
4081 );
4082 add_raw("model.layers.1.mlp.up_proj.weight", vec![inter, hidden], &g);
4083 add_raw(
4084 "model.layers.1.mlp.down_proj.weight",
4085 vec![hidden, inter],
4086 &d,
4087 );
4088 add_raw("model.norm.weight", vec![hidden], &n1);
4089 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
4090
4091 let cfg = json!({
4092 "format": "aria-quant-bundle",
4093 "format_version": 2,
4094 "quantization": "test",
4095 "group_size_default": 32,
4096 "hadamard_seed": 0,
4097 "model": {
4098 "hidden_size": hidden,
4099 "num_layers": layers,
4100 "num_attention_heads": n_heads,
4101 "num_kv_heads": n_kv,
4102 "head_dim": head_dim,
4103 "intermediate_size": inter,
4104 "vocab_size": vocab,
4105 "context_length": 32,
4106 "rope_theta": 10000.0,
4107 "conv_l_cache": kernel,
4108 "layer_types": ["conv", "full_attention"]
4109 },
4110 "tensors": tensors
4111 });
4112 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
4113 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
4114
4115 let mut s = SessionBuilder::new()
4116 .model(dir.path())
4117 .family("lfm/lfm2-350m")
4118 .build()
4119 .unwrap();
4120 let gen = s
4121 .generate(
4122 &[1, 2, 3],
4123 &GenerateOpts {
4124 max_tokens: 2,
4125 temperature: 0.0,
4126 },
4127 )
4128 .unwrap();
4129 assert_eq!(gen.tokens.len(), 2);
4130 }
4131
4132 #[test]
4133 fn moe_topk_experts_generate() {
4134 let dir = tempfile::tempdir().unwrap();
4135 let hidden = 8usize;
4136 let layers = 1usize;
4137 let inter = 16usize;
4138 let vocab = 16usize;
4139 let n_heads = 2usize;
4140 let n_kv = 1usize;
4141 let head_dim = 4usize;
4142 let q_dim = n_heads * head_dim;
4143 let k_dim = n_kv * head_dim;
4144 let n_experts = 4usize;
4145 let top_k = 2usize;
4146
4147 let mut tensors = serde_json::Map::new();
4148 let mut bin = Vec::new();
4149 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
4150 let offset = bin.len();
4151 for &v in data {
4152 bin.extend_from_slice(&v.to_le_bytes());
4153 }
4154 let nbytes = data.len() * 4;
4155 let mut meta = serde_json::Map::new();
4156 meta.insert("kind".into(), json!("raw"));
4157 meta.insert("dtype".into(), json!("f32"));
4158 meta.insert("shape".into(), json!(shape));
4159 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
4160 tensors.insert(name.to_string(), Value::Object(meta));
4161 };
4162 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
4163 add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
4164 let n1 = vec![1.0f32; hidden];
4165 add_raw("model.layers.0.input_layernorm.weight", vec![hidden], &n1);
4166 add_raw(
4167 "model.layers.0.post_attention_layernorm.weight",
4168 vec![hidden],
4169 &n1,
4170 );
4171 let wq = vec![0.02f32; q_dim * hidden];
4172 let wk = vec![0.02f32; k_dim * hidden];
4173 let wv = vec![0.02f32; k_dim * hidden];
4174 let wo = vec![0.02f32; hidden * q_dim];
4175 add_raw(
4176 "model.layers.0.self_attn.q_proj.weight",
4177 vec![q_dim, hidden],
4178 &wq,
4179 );
4180 add_raw(
4181 "model.layers.0.self_attn.k_proj.weight",
4182 vec![k_dim, hidden],
4183 &wk,
4184 );
4185 add_raw(
4186 "model.layers.0.self_attn.v_proj.weight",
4187 vec![k_dim, hidden],
4188 &wv,
4189 );
4190 add_raw(
4191 "model.layers.0.self_attn.o_proj.weight",
4192 vec![hidden, q_dim],
4193 &wo,
4194 );
4195 let router: Vec<f32> = (0..n_experts * hidden)
4196 .map(|i| ((i % n_experts) as f32) * 0.1)
4197 .collect();
4198 add_raw(
4199 "model.layers.0.block_sparse_moe.gate.weight",
4200 vec![n_experts, hidden],
4201 &router,
4202 );
4203 let g = vec![0.02f32; inter * hidden];
4204 let d = vec![0.02f32; hidden * inter];
4205 for e in 0..n_experts {
4206 add_raw(
4207 &format!("model.layers.0.block_sparse_moe.experts.{e}.w1.weight"),
4208 vec![inter, hidden],
4209 &g,
4210 );
4211 add_raw(
4212 &format!("model.layers.0.block_sparse_moe.experts.{e}.w3.weight"),
4213 vec![inter, hidden],
4214 &g,
4215 );
4216 add_raw(
4217 &format!("model.layers.0.block_sparse_moe.experts.{e}.w2.weight"),
4218 vec![hidden, inter],
4219 &d,
4220 );
4221 }
4222 add_raw("model.norm.weight", vec![hidden], &n1);
4223 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
4224
4225 let cfg = json!({
4226 "format": "aria-quant-bundle",
4227 "format_version": 2,
4228 "quantization": "test",
4229 "group_size_default": 32,
4230 "hadamard_seed": 0,
4231 "model": {
4232 "hidden_size": hidden,
4233 "num_layers": layers,
4234 "num_attention_heads": n_heads,
4235 "num_kv_heads": n_kv,
4236 "head_dim": head_dim,
4237 "intermediate_size": inter,
4238 "vocab_size": vocab,
4239 "context_length": 32,
4240 "rope_theta": 10000.0,
4241 "num_experts": n_experts,
4242 "num_experts_per_tok": top_k
4243 },
4244 "tensors": tensors
4245 });
4246 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
4247 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
4248
4249 let mut s = SessionBuilder::new()
4250 .model(dir.path())
4251 .family("inkling/inkling-small")
4252 .build()
4253 .unwrap();
4254 assert_eq!(s.arch(), ArchClass::TextMoE);
4255 assert_eq!(s.graph_hook_name(), "text_moe_decoder");
4256 let gen = s
4257 .generate(
4258 &[1, 2],
4259 &GenerateOpts {
4260 max_tokens: 2,
4261 temperature: 0.0,
4262 },
4263 )
4264 .unwrap();
4265 assert_eq!(gen.tokens.len(), 2);
4266 }
4267
4268 #[test]
4269 fn tiny_q4_codebook_weights_unrotate_on_load() {
4270 let dir = tempfile::tempdir().unwrap();
4271 write_tiny_q4_bundle(dir.path()).unwrap();
4272 let b = load_bundle(dir.path()).unwrap();
4273 let w = b.weight_loaded("blk.0.attn_q.weight").unwrap();
4274 assert!(
4275 w.hdm_seed.is_none(),
4276 "reconstruct_weight path stores original-space W for linear()"
4277 );
4278 let mut s = SessionBuilder::new()
4279 .model(dir.path())
4280 .family("gemma/gemma-4-e2b-it")
4281 .build()
4282 .unwrap();
4283 let gen = s
4284 .generate(
4285 &[1, 2],
4286 &GenerateOpts {
4287 max_tokens: 2,
4288 temperature: 0.0,
4289 },
4290 )
4291 .unwrap();
4292 assert_eq!(gen.tokens.len(), 2);
4293 }
4294
4295 #[test]
4296 fn gemma_hidden_act_geglu_and_qk_norm() {
4297 let dir = tempfile::tempdir().unwrap();
4298 let hidden = 8usize;
4299 let layers = 1usize;
4300 let inter = 16usize;
4301 let vocab = 16usize;
4302 let n_heads = 2usize;
4303 let n_kv = 1usize;
4304 let head_dim = 4usize;
4305 let q_dim = n_heads * head_dim;
4306 let k_dim = n_kv * head_dim;
4307 let p = "model.language_model";
4308
4309 let mut tensors = serde_json::Map::new();
4310 let mut bin = Vec::new();
4311 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
4312 let offset = bin.len();
4313 for &v in data {
4314 bin.extend_from_slice(&v.to_le_bytes());
4315 }
4316 let nbytes = data.len() * 4;
4317 let mut meta = serde_json::Map::new();
4318 meta.insert("kind".into(), json!("raw"));
4319 meta.insert("dtype".into(), json!("f32"));
4320 meta.insert("shape".into(), json!(shape));
4321 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
4322 tensors.insert(name.to_string(), Value::Object(meta));
4323 };
4324 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
4325 add_raw(
4326 &format!("{p}.embed_tokens.weight"),
4327 vec![vocab, hidden],
4328 &emb,
4329 );
4330 let n1 = vec![1.0f32; hidden];
4331 let qn = vec![1.0f32; head_dim];
4332 let kn = vec![1.0f32; head_dim];
4333 let wq = vec![0.01f32; q_dim * hidden];
4334 let wk = vec![0.01f32; k_dim * hidden];
4335 let wv = vec![0.01f32; k_dim * hidden];
4336 let wo = vec![0.01f32; hidden * q_dim];
4337 add_raw(
4338 &format!("{p}.layers.0.input_layernorm.weight"),
4339 vec![hidden],
4340 &n1,
4341 );
4342 add_raw(
4343 &format!("{p}.layers.0.pre_feedforward_layernorm.weight"),
4344 vec![hidden],
4345 &n1,
4346 );
4347 add_raw(
4348 &format!("{p}.layers.0.self_attn.q_proj.weight"),
4349 vec![q_dim, hidden],
4350 &wq,
4351 );
4352 add_raw(
4353 &format!("{p}.layers.0.self_attn.k_proj.weight"),
4354 vec![k_dim, hidden],
4355 &wk,
4356 );
4357 add_raw(
4358 &format!("{p}.layers.0.self_attn.v_proj.weight"),
4359 vec![k_dim, hidden],
4360 &wv,
4361 );
4362 add_raw(
4363 &format!("{p}.layers.0.self_attn.o_proj.weight"),
4364 vec![hidden, q_dim],
4365 &wo,
4366 );
4367 add_raw(
4368 &format!("{p}.layers.0.self_attn.q_norm.weight"),
4369 vec![head_dim],
4370 &qn,
4371 );
4372 add_raw(
4373 &format!("{p}.layers.0.self_attn.k_norm.weight"),
4374 vec![head_dim],
4375 &kn,
4376 );
4377 let g = vec![0.01f32; inter * hidden];
4378 let d = vec![0.01f32; hidden * inter];
4379 add_raw(
4380 &format!("{p}.layers.0.mlp.gate_proj.weight"),
4381 vec![inter, hidden],
4382 &g,
4383 );
4384 add_raw(
4385 &format!("{p}.layers.0.mlp.up_proj.weight"),
4386 vec![inter, hidden],
4387 &g,
4388 );
4389 add_raw(
4390 &format!("{p}.layers.0.mlp.down_proj.weight"),
4391 vec![hidden, inter],
4392 &d,
4393 );
4394 add_raw(&format!("{p}.norm.weight"), vec![hidden], &n1);
4395 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
4396
4397 let cfg = json!({
4398 "format": "aria-quant-bundle",
4399 "format_version": 2,
4400 "quantization": "test",
4401 "group_size_default": 32,
4402 "hadamard_seed": 0,
4403 "model": {
4404 "hidden_size": hidden,
4405 "num_layers": layers,
4406 "num_attention_heads": n_heads,
4407 "num_kv_heads": n_kv,
4408 "head_dim": head_dim,
4409 "global_head_dim": head_dim,
4410 "sliding_window": 512,
4411 "partial_rotary_factor": 0.25,
4412 "intermediate_size": inter,
4413 "vocab_size": vocab,
4414 "context_length": 32,
4415 "rope_theta": 10000.0,
4416 "hidden_act": "gelu_pytorch_tanh",
4417 "layer_types": ["full_attention"]
4418 },
4419 "tensors": tensors
4420 });
4421 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
4422 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
4423
4424 let mut s = SessionBuilder::new()
4425 .model(dir.path())
4426 .family("gemma/gemma-4-e2b-it")
4427 .build()
4428 .unwrap();
4429 assert_eq!(s.config().hidden_act.as_deref(), Some("gelu_pytorch_tanh"));
4430 let gen = s
4431 .generate(
4432 &[1, 2],
4433 &GenerateOpts {
4434 max_tokens: 2,
4435 temperature: 0.0,
4436 },
4437 )
4438 .unwrap();
4439 assert_eq!(gen.tokens.len(), 2);
4440 }
4441
4442 #[test]
4443 fn gated_deltanet_and_full_attn_generate() {
4444 let dir = tempfile::tempdir().unwrap();
4445 let hidden = 8usize;
4446 let inter = 16usize;
4447 let vocab = 16usize;
4448 let n_heads = 2usize;
4449 let n_kv = 1usize;
4450 let head_dim = 4usize;
4451 let q_dim = n_heads * head_dim;
4452 let k_dim = n_kv * head_dim;
4453 let n_lin = 2usize;
4454 let hk = 4usize;
4455 let hv = 4usize;
4456 let key_dim = n_lin * hk;
4457 let value_dim = n_lin * hv;
4458 let conv_k = 4usize;
4459 let conv_dim = key_dim * 2 + value_dim;
4460
4461 let mut tensors = serde_json::Map::new();
4462 let mut bin = Vec::new();
4463 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
4464 let offset = bin.len();
4465 for &v in data {
4466 bin.extend_from_slice(&v.to_le_bytes());
4467 }
4468 let nbytes = data.len() * 4;
4469 let mut meta = serde_json::Map::new();
4470 meta.insert("kind".into(), json!("raw"));
4471 meta.insert("dtype".into(), json!("f32"));
4472 meta.insert("shape".into(), json!(shape));
4473 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
4474 tensors.insert(name.to_string(), Value::Object(meta));
4475 };
4476 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
4477 add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
4478 let n1 = vec![1.0f32; hidden];
4479 add_raw("model.layers.0.input_layernorm.weight", vec![hidden], &n1);
4480 add_raw(
4481 "model.layers.0.post_attention_layernorm.weight",
4482 vec![hidden],
4483 &n1,
4484 );
4485 let qkvz = vec![0.02f32; (2 * key_dim + 2 * value_dim) * hidden];
4486 let ba = vec![0.1f32; 2 * n_lin * hidden];
4487 let conv = vec![0.05f32; conv_dim * conv_k];
4488 let a_log = vec![0.5f32; n_lin];
4489 let dt = vec![1.0f32; n_lin];
4490 let outp = vec![0.02f32; hidden * value_dim];
4491 add_raw(
4492 "model.layers.0.linear_attn.in_proj_qkvz.weight",
4493 vec![2 * key_dim + 2 * value_dim, hidden],
4494 &qkvz,
4495 );
4496 add_raw(
4497 "model.layers.0.linear_attn.in_proj_ba.weight",
4498 vec![2 * n_lin, hidden],
4499 &ba,
4500 );
4501 add_raw(
4502 "model.layers.0.linear_attn.conv1d.weight",
4503 vec![conv_dim, conv_k],
4504 &conv,
4505 );
4506 add_raw("model.layers.0.linear_attn.A_log", vec![n_lin], &a_log);
4507 add_raw("model.layers.0.linear_attn.dt_bias", vec![n_lin], &dt);
4508 add_raw(
4509 "model.layers.0.linear_attn.out_proj.weight",
4510 vec![hidden, value_dim],
4511 &outp,
4512 );
4513 let g = vec![0.02f32; inter * hidden];
4514 let d = vec![0.02f32; hidden * inter];
4515 add_raw(
4516 "model.layers.0.mlp.gate_proj.weight",
4517 vec![inter, hidden],
4518 &g,
4519 );
4520 add_raw("model.layers.0.mlp.up_proj.weight", vec![inter, hidden], &g);
4521 add_raw(
4522 "model.layers.0.mlp.down_proj.weight",
4523 vec![hidden, inter],
4524 &d,
4525 );
4526
4527 add_raw("model.layers.1.input_layernorm.weight", vec![hidden], &n1);
4528 add_raw(
4529 "model.layers.1.post_attention_layernorm.weight",
4530 vec![hidden],
4531 &n1,
4532 );
4533 let wq = vec![0.02f32; q_dim * hidden];
4534 let wk = vec![0.02f32; k_dim * hidden];
4535 let wv = vec![0.02f32; k_dim * hidden];
4536 let wo = vec![0.02f32; hidden * q_dim];
4537 add_raw(
4538 "model.layers.1.self_attn.q_proj.weight",
4539 vec![q_dim, hidden],
4540 &wq,
4541 );
4542 add_raw(
4543 "model.layers.1.self_attn.k_proj.weight",
4544 vec![k_dim, hidden],
4545 &wk,
4546 );
4547 add_raw(
4548 "model.layers.1.self_attn.v_proj.weight",
4549 vec![k_dim, hidden],
4550 &wv,
4551 );
4552 add_raw(
4553 "model.layers.1.self_attn.o_proj.weight",
4554 vec![hidden, q_dim],
4555 &wo,
4556 );
4557 add_raw(
4558 "model.layers.1.mlp.gate_proj.weight",
4559 vec![inter, hidden],
4560 &g,
4561 );
4562 add_raw("model.layers.1.mlp.up_proj.weight", vec![inter, hidden], &g);
4563 add_raw(
4564 "model.layers.1.mlp.down_proj.weight",
4565 vec![hidden, inter],
4566 &d,
4567 );
4568 add_raw("model.norm.weight", vec![hidden], &n1);
4569 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
4570
4571 let cfg = json!({
4572 "format": "aria-quant-bundle",
4573 "format_version": 2,
4574 "quantization": "test",
4575 "hadamard_seed": 0,
4576 "model": {
4577 "hidden_size": hidden,
4578 "num_layers": 2,
4579 "num_attention_heads": n_heads,
4580 "num_kv_heads": n_kv,
4581 "head_dim": head_dim,
4582 "intermediate_size": inter,
4583 "vocab_size": vocab,
4584 "context_length": 32,
4585 "rope_theta": 10000.0,
4586 "layer_types": ["linear_attention", "full_attention"]
4587 },
4588 "tensors": tensors
4589 });
4590 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
4591 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
4592 let mut s = SessionBuilder::new()
4593 .model(dir.path())
4594 .family("qwen/qwen3.5-2b")
4595 .build()
4596 .unwrap();
4597 assert_eq!(s.config().partial_rotary_factor, Some(0.25));
4598 assert!(
4599 (s.config().rope_theta - 10_000_000.0).abs() < 1.0,
4600 "Qwen3.5 Llama-default rope_theta must become 1e7, got {}",
4601 s.config().rope_theta
4602 );
4603 let gen = s
4604 .generate(
4605 &[1, 2, 3],
4606 &GenerateOpts {
4607 max_tokens: 2,
4608 temperature: 0.0,
4609 },
4610 )
4611 .unwrap();
4612 assert_eq!(gen.tokens.len(), 2);
4613 }
4614
4615 #[test]
4616 fn gated_deltanet_split_qwen35_projections_generate() {
4617 let dir = tempfile::tempdir().unwrap();
4618 let hidden = 8usize;
4619 let inter = 16usize;
4620 let vocab = 16usize;
4621 let n_heads = 2usize;
4622 let n_kv = 1usize;
4623 let head_dim = 4usize;
4624 let q_dim = n_heads * head_dim;
4625 let k_dim = n_kv * head_dim;
4626 let n_lin = 2usize;
4627 let hk = 4usize;
4628 let hv = 4usize;
4629 let key_dim = n_lin * hk;
4630 let value_dim = n_lin * hv;
4631 let conv_k = 4usize;
4632 let conv_dim = key_dim * 2 + value_dim;
4633
4634 let mut tensors = serde_json::Map::new();
4635 let mut bin = Vec::new();
4636 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
4637 let offset = bin.len();
4638 for &v in data {
4639 bin.extend_from_slice(&v.to_le_bytes());
4640 }
4641 let nbytes = data.len() * 4;
4642 let mut meta = serde_json::Map::new();
4643 meta.insert("kind".into(), json!("raw"));
4644 meta.insert("dtype".into(), json!("f32"));
4645 meta.insert("shape".into(), json!(shape));
4646 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
4647 tensors.insert(name.to_string(), Value::Object(meta));
4648 };
4649 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
4650 add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
4651 let n1 = vec![1.0f32; hidden];
4652 add_raw("model.layers.0.input_layernorm.weight", vec![hidden], &n1);
4653 add_raw(
4654 "model.layers.0.post_attention_layernorm.weight",
4655 vec![hidden],
4656 &n1,
4657 );
4658 let qkv = vec![0.02f32; (2 * key_dim + value_dim) * hidden];
4659 let z = vec![0.02f32; value_dim * hidden];
4660 let proj_b = vec![0.1f32; n_lin * hidden];
4661 let proj_a = vec![0.1f32; n_lin * hidden];
4662 let conv = vec![0.05f32; conv_dim * conv_k];
4663 let a_log = vec![0.5f32; n_lin];
4664 let dt = vec![1.0f32; n_lin];
4665 let outp = vec![0.02f32; hidden * value_dim];
4666 add_raw(
4667 "model.layers.0.linear_attn.in_proj_qkv.weight",
4668 vec![2 * key_dim + value_dim, hidden],
4669 &qkv,
4670 );
4671 add_raw(
4672 "model.layers.0.linear_attn.in_proj_z.weight",
4673 vec![value_dim, hidden],
4674 &z,
4675 );
4676 add_raw(
4677 "model.layers.0.linear_attn.in_proj_b.weight",
4678 vec![n_lin, hidden],
4679 &proj_b,
4680 );
4681 add_raw(
4682 "model.layers.0.linear_attn.in_proj_a.weight",
4683 vec![n_lin, hidden],
4684 &proj_a,
4685 );
4686 add_raw(
4687 "model.layers.0.linear_attn.conv1d.weight",
4688 vec![conv_dim, conv_k],
4689 &conv,
4690 );
4691 add_raw("model.layers.0.linear_attn.A_log", vec![n_lin], &a_log);
4692 add_raw("model.layers.0.linear_attn.dt_bias", vec![n_lin], &dt);
4693 add_raw(
4694 "model.layers.0.linear_attn.out_proj.weight",
4695 vec![hidden, value_dim],
4696 &outp,
4697 );
4698 let g = vec![0.02f32; inter * hidden];
4699 let d = vec![0.02f32; hidden * inter];
4700 add_raw(
4701 "model.layers.0.mlp.gate_proj.weight",
4702 vec![inter, hidden],
4703 &g,
4704 );
4705 add_raw("model.layers.0.mlp.up_proj.weight", vec![inter, hidden], &g);
4706 add_raw(
4707 "model.layers.0.mlp.down_proj.weight",
4708 vec![hidden, inter],
4709 &d,
4710 );
4711
4712 add_raw("model.layers.1.input_layernorm.weight", vec![hidden], &n1);
4713 add_raw(
4714 "model.layers.1.post_attention_layernorm.weight",
4715 vec![hidden],
4716 &n1,
4717 );
4718 let wq = vec![0.02f32; q_dim * hidden];
4719 let wk = vec![0.02f32; k_dim * hidden];
4720 let wv = vec![0.02f32; k_dim * hidden];
4721 let wo = vec![0.02f32; hidden * q_dim];
4722 add_raw(
4723 "model.layers.1.self_attn.q_proj.weight",
4724 vec![q_dim, hidden],
4725 &wq,
4726 );
4727 add_raw(
4728 "model.layers.1.self_attn.k_proj.weight",
4729 vec![k_dim, hidden],
4730 &wk,
4731 );
4732 add_raw(
4733 "model.layers.1.self_attn.v_proj.weight",
4734 vec![k_dim, hidden],
4735 &wv,
4736 );
4737 add_raw(
4738 "model.layers.1.self_attn.o_proj.weight",
4739 vec![hidden, q_dim],
4740 &wo,
4741 );
4742 add_raw(
4743 "model.layers.1.mlp.gate_proj.weight",
4744 vec![inter, hidden],
4745 &g,
4746 );
4747 add_raw("model.layers.1.mlp.up_proj.weight", vec![inter, hidden], &g);
4748 add_raw(
4749 "model.layers.1.mlp.down_proj.weight",
4750 vec![hidden, inter],
4751 &d,
4752 );
4753 add_raw("model.norm.weight", vec![hidden], &n1);
4754 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
4755
4756 let cfg = json!({
4757 "format": "aria-quant-bundle",
4758 "format_version": 2,
4759 "quantization": "test",
4760 "hadamard_seed": 0,
4761 "model": {
4762 "hidden_size": hidden,
4763 "num_layers": 2,
4764 "num_attention_heads": n_heads,
4765 "num_kv_heads": n_kv,
4766 "head_dim": head_dim,
4767 "intermediate_size": inter,
4768 "vocab_size": vocab,
4769 "context_length": 32,
4770 "rope_theta": 10000.0,
4771 "layer_types": ["linear_attention", "full_attention"]
4772 },
4773 "tensors": tensors
4774 });
4775 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
4776 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
4777 let mut s = SessionBuilder::new()
4778 .model(dir.path())
4779 .family("qwen/qwen3.5-0.8b")
4780 .build()
4781 .unwrap();
4782 let gen = s
4783 .generate(
4784 &[1, 2, 3],
4785 &GenerateOpts {
4786 max_tokens: 2,
4787 temperature: 0.0,
4788 },
4789 )
4790 .unwrap();
4791 assert_eq!(gen.tokens.len(), 2);
4792 }
4793
4794 #[test]
4795 fn vision_and_action_consume_bundle_weights() {
4796 let dir = tempfile::tempdir().unwrap();
4797 write_tiny_q4_bundle(dir.path()).unwrap();
4798 let cfg_path = dir.path().join("config.json");
4799 let mut cfg: Value =
4800 serde_json::from_str(&std::fs::read_to_string(&cfg_path).unwrap()).unwrap();
4801 let hidden = cfg["model"]["hidden_size"].as_u64().unwrap() as usize;
4802 let mut tensors = cfg["tensors"].as_object().cloned().unwrap();
4803 let mut bin = std::fs::read(dir.path().join("weight.bin")).unwrap();
4804 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
4805 let offset = bin.len();
4806 for &v in data {
4807 bin.extend_from_slice(&v.to_le_bytes());
4808 }
4809 let nbytes = data.len() * 4;
4810 tensors.insert(
4811 name.to_string(),
4812 json!({
4813 "kind": "raw",
4814 "dtype": "f32",
4815 "shape": shape,
4816 "offsets": { "data": [offset, nbytes] }
4817 }),
4818 );
4819 };
4820 let vis = vec![0.1f32; hidden * 3];
4821 add_raw("mm_projector.weight", vec![hidden, 3], &vis);
4822 let act_dim = 7usize;
4823 let act = vec![0.05f32; act_dim * hidden];
4824 add_raw("action_head.weight", vec![act_dim, hidden], &act);
4825 cfg["tensors"] = Value::Object(tensors);
4826 std::fs::write(&cfg_path, cfg.to_string()).unwrap();
4827 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
4828
4829 let s = SessionBuilder::new()
4830 .model(dir.path())
4831 .family("lfm/lfm2-vl-450m")
4832 .build()
4833 .unwrap();
4834 let rgb = vec![10u8; 3 * 4 * 4];
4835 let pref = s.vision_prefix(&rgb, 4, 4).unwrap();
4836 assert_eq!(pref.len(), hidden);
4837
4838 let vla = SessionBuilder::new()
4839 .model(dir.path())
4840 .family("openvla/openvla-7b")
4841 .build()
4842 .unwrap();
4843 let a = vla.predict_action("move", act_dim).unwrap();
4844 assert_eq!(a.len(), act_dim);
4845 }
4846
4847 #[test]
4848 fn load_real_hf_named_bundle_if_present() {
4849 let Ok(path) = std::env::var("ARIA_SMOKE_BUNDLE") else {
4851 return;
4852 };
4853 let path = std::path::Path::new(&path);
4854 if !path.join("config.json").is_file() {
4855 return;
4856 }
4857 let family = if path.to_string_lossy().contains("gemma-4") {
4858 "gemma/gemma-4-e2b-it"
4859 } else {
4860 "qwen/qwen3-0.6b"
4861 };
4862 let s = SessionBuilder::new()
4863 .model(path)
4864 .family(family)
4865 .build()
4866 .unwrap_or_else(|e| panic!("{family} bundle should materialize: {e}"));
4867 assert!(s.config().num_layers > 0);
4868 assert!(s.config().hidden_size > 0);
4869 if family.contains("gemma-4") && s.config().hidden_size >= 1024 {
4870 assert!(
4871 s.weights.ple.is_some(),
4872 "real Gemma-4 q4 must load codebook PLE"
4873 );
4874 let hidden = s.config().hidden_size;
4875 let vocab = s.config().vocab_size;
4876 assert!(
4877 s.weights.emb.data.len() >= vocab.saturating_mul(hidden),
4878 "embed table too small for vocab={vocab} hidden={hidden}"
4879 );
4880 }
4881 }
4882
4883 #[test]
4884 fn gemma4_four_norm_ple_and_tied_embed_generate() {
4885 let dir = tempfile::tempdir().unwrap();
4886 let hidden = 8usize;
4887 let layers = 1usize;
4888 let inter = 16usize;
4889 let vocab = 16usize;
4890 let n_heads = 2usize;
4891 let n_kv = 1usize;
4892 let head_dim = 4usize;
4893 let q_dim = n_heads * head_dim;
4894 let k_dim = n_kv * head_dim;
4895 let ple_d = 4usize;
4896 let p = "model.language_model";
4897
4898 let mut tensors = serde_json::Map::new();
4899 let mut bin = Vec::new();
4900 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
4901 let offset = bin.len();
4902 for &v in data {
4903 bin.extend_from_slice(&v.to_le_bytes());
4904 }
4905 let nbytes = data.len() * 4;
4906 let mut meta = serde_json::Map::new();
4907 meta.insert("kind".into(), json!("raw"));
4908 meta.insert("dtype".into(), json!("f32"));
4909 meta.insert("shape".into(), json!(shape));
4910 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
4911 tensors.insert(name.to_string(), Value::Object(meta));
4912 };
4913 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
4914 add_raw(
4915 &format!("{p}.embed_tokens.weight"),
4916 vec![vocab, hidden],
4917 &emb,
4918 );
4919 let n1 = vec![1.0f32; hidden];
4920 add_raw(
4921 &format!("{p}.layers.0.input_layernorm.weight"),
4922 vec![hidden],
4923 &n1,
4924 );
4925 add_raw(
4926 &format!("{p}.layers.0.post_attention_layernorm.weight"),
4927 vec![hidden],
4928 &n1,
4929 );
4930 add_raw(
4931 &format!("{p}.layers.0.pre_feedforward_layernorm.weight"),
4932 vec![hidden],
4933 &n1,
4934 );
4935 add_raw(
4936 &format!("{p}.layers.0.post_feedforward_layernorm.weight"),
4937 vec![hidden],
4938 &n1,
4939 );
4940 add_raw(&format!("{p}.layers.0.layer_scalar"), vec![1], &[0.5f32]);
4942 let wq = vec![0.01f32; q_dim * hidden];
4943 let wk = vec![0.01f32; k_dim * hidden];
4944 let wv = vec![0.01f32; k_dim * hidden];
4945 let wo = vec![0.01f32; hidden * q_dim];
4946 add_raw(
4947 &format!("{p}.layers.0.self_attn.q_proj.weight"),
4948 vec![q_dim, hidden],
4949 &wq,
4950 );
4951 add_raw(
4952 &format!("{p}.layers.0.self_attn.k_proj.weight"),
4953 vec![k_dim, hidden],
4954 &wk,
4955 );
4956 add_raw(
4957 &format!("{p}.layers.0.self_attn.v_proj.weight"),
4958 vec![k_dim, hidden],
4959 &wv,
4960 );
4961 add_raw(
4962 &format!("{p}.layers.0.self_attn.o_proj.weight"),
4963 vec![hidden, q_dim],
4964 &wo,
4965 );
4966 let g = vec![0.01f32; inter * hidden];
4967 let d = vec![0.01f32; hidden * inter];
4968 add_raw(
4969 &format!("{p}.layers.0.mlp.gate_proj.weight"),
4970 vec![inter, hidden],
4971 &g,
4972 );
4973 add_raw(
4974 &format!("{p}.layers.0.mlp.up_proj.weight"),
4975 vec![inter, hidden],
4976 &g,
4977 );
4978 add_raw(
4979 &format!("{p}.layers.0.mlp.down_proj.weight"),
4980 vec![hidden, inter],
4981 &d,
4982 );
4983 let packed = layers * ple_d;
4984 let ple_emb = vec![0.02f32; vocab * packed];
4985 add_raw(
4986 &format!("{p}.embed_tokens_per_layer.weight"),
4987 vec![vocab, packed],
4988 &ple_emb,
4989 );
4990 let ple_proj = vec![0.01f32; packed * hidden];
4991 add_raw(
4992 &format!("{p}.per_layer_model_projection.weight"),
4993 vec![packed, hidden],
4994 &ple_proj,
4995 );
4996 let ple_pn = vec![1.0f32; ple_d];
4997 add_raw(
4998 &format!("{p}.per_layer_projection_norm.weight"),
4999 vec![ple_d],
5000 &ple_pn,
5001 );
5002 let ple_gate = vec![0.01f32; ple_d * hidden];
5003 let ple_out = vec![0.01f32; hidden * ple_d];
5004 add_raw(
5005 &format!("{p}.layers.0.per_layer_input_gate.weight"),
5006 vec![ple_d, hidden],
5007 &ple_gate,
5008 );
5009 add_raw(
5010 &format!("{p}.layers.0.per_layer_projection.weight"),
5011 vec![hidden, ple_d],
5012 &ple_out,
5013 );
5014 add_raw(
5015 &format!("{p}.layers.0.post_per_layer_input_norm.weight"),
5016 vec![hidden],
5017 &n1,
5018 );
5019 add_raw(&format!("{p}.norm.weight"), vec![hidden], &n1);
5020 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
5021
5022 let cfg = json!({
5023 "format": "aria-quant-bundle",
5024 "format_version": 2,
5025 "quantization": "test",
5026 "group_size_default": 32,
5027 "hadamard_seed": 0,
5028 "model": {
5029 "hidden_size": hidden,
5030 "num_layers": layers,
5031 "num_attention_heads": n_heads,
5032 "num_kv_heads": n_kv,
5033 "intermediate_size": inter,
5034 "vocab_size": vocab,
5035 "context_length": 32,
5036 "rope_theta": 10000.0,
5037 "hidden_act": "gelu_pytorch_tanh",
5038 "tie_word_embeddings": true,
5039 "head_dim": head_dim,
5040 "global_head_dim": head_dim,
5041 "sliding_window": 512,
5042 "partial_rotary_factor": 0.25,
5043 "layer_types": ["full_attention"]
5044 },
5045 "tensors": tensors
5046 });
5047 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
5048 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
5049
5050 let mut s = SessionBuilder::new()
5051 .model(dir.path())
5052 .family("gemma/gemma-4-e2b-it")
5053 .build()
5054 .unwrap();
5055 assert!((s.embed_scale - (hidden as f32).sqrt()).abs() < 1e-5);
5056 assert!(s.weights.ple.is_some());
5057 assert!((s.weights.layers[0].layer_scalar - 0.5).abs() < 1e-6);
5058 assert!(s.weights.layers[0].post_attn_norm.is_some());
5059 assert!(s.weights.layers[0].post_ffn_norm.is_some());
5060 let prompt = vec![1u32, 2];
5061 let batched = s
5062 .generate(
5063 &prompt,
5064 &GenerateOpts {
5065 max_tokens: 3,
5066 temperature: 0.0,
5067 },
5068 )
5069 .unwrap();
5070 let step = s
5071 .generate(
5072 &prompt,
5073 &GenerateOpts {
5074 max_tokens: 3,
5075 temperature: 0.0,
5076 },
5077 )
5078 .unwrap();
5079 assert_eq!(batched.tokens, step.tokens);
5080 assert_eq!(batched.tokens.len(), 3);
5081 assert_eq!(s.config().sliding_window, Some(512));
5082 }
5083
5084 #[test]
5085 fn gemma4_ple_required_gate() {
5086 assert!(!gemma4_requires_ple("gemma/gemma-4-e2b-it", 64));
5087 assert!(gemma4_requires_ple("gemma/gemma-4-e2b-it", 1024));
5088 assert!(gemma4_requires_ple("gemma/gemma-4-e2b-it", 1536));
5089 assert!(gemma4_requires_ple("gemma/gemma-3n-e2b-it", 2048));
5090 assert!(!gemma4_requires_ple("qwen/qwen3-0.6b", 1536));
5091 assert!(!gemma3n_requires_altup("gemma/gemma-3n-e2b-it", 64));
5092 assert!(gemma3n_requires_altup("gemma/gemma-3n-e2b-it", 2048));
5093 assert!(!gemma3n_requires_altup("gemma/gemma-4-e2b-it", 2048));
5094 }
5095
5096 #[test]
5097 fn gemma3n_altup_laurel_ple_generate() {
5098 let dir = tempfile::tempdir().unwrap();
5099 let hidden = 8usize;
5100 let layers = 1usize;
5101 let inter = 16usize;
5102 let vocab = 16usize;
5103 let n_heads = 2usize;
5104 let n_kv = 1usize;
5105 let head_dim = 4usize;
5106 let q_dim = n_heads * head_dim;
5107 let k_dim = n_kv * head_dim;
5108 let ple_d = 4usize;
5109 let rank = 2usize;
5110 let n_alt = 4usize;
5111 let p = "model.language_model";
5112
5113 let mut tensors = serde_json::Map::new();
5114 let mut bin = Vec::new();
5115 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
5116 let offset = bin.len();
5117 for &v in data {
5118 bin.extend_from_slice(&v.to_le_bytes());
5119 }
5120 let nbytes = data.len() * 4;
5121 let mut meta = serde_json::Map::new();
5122 meta.insert("kind".into(), json!("raw"));
5123 meta.insert("dtype".into(), json!("f32"));
5124 meta.insert("shape".into(), json!(shape));
5125 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
5126 tensors.insert(name.to_string(), Value::Object(meta));
5127 };
5128 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
5129 add_raw(
5130 &format!("{p}.embed_tokens.weight"),
5131 vec![vocab, hidden],
5132 &emb,
5133 );
5134 let n1 = vec![1.0f32; hidden];
5135 add_raw(
5136 &format!("{p}.layers.0.input_layernorm.weight"),
5137 vec![hidden],
5138 &n1,
5139 );
5140 add_raw(
5141 &format!("{p}.layers.0.post_attention_layernorm.weight"),
5142 vec![hidden],
5143 &n1,
5144 );
5145 add_raw(
5146 &format!("{p}.layers.0.pre_feedforward_layernorm.weight"),
5147 vec![hidden],
5148 &n1,
5149 );
5150 add_raw(
5151 &format!("{p}.layers.0.post_feedforward_layernorm.weight"),
5152 vec![hidden],
5153 &n1,
5154 );
5155 let wq = vec![0.01f32; q_dim * hidden];
5156 let wk = vec![0.01f32; k_dim * hidden];
5157 let wv = vec![0.01f32; k_dim * hidden];
5158 let wo = vec![0.01f32; hidden * q_dim];
5159 add_raw(
5160 &format!("{p}.layers.0.self_attn.q_proj.weight"),
5161 vec![q_dim, hidden],
5162 &wq,
5163 );
5164 add_raw(
5165 &format!("{p}.layers.0.self_attn.k_proj.weight"),
5166 vec![k_dim, hidden],
5167 &wk,
5168 );
5169 add_raw(
5170 &format!("{p}.layers.0.self_attn.v_proj.weight"),
5171 vec![k_dim, hidden],
5172 &wv,
5173 );
5174 add_raw(
5175 &format!("{p}.layers.0.self_attn.o_proj.weight"),
5176 vec![hidden, q_dim],
5177 &wo,
5178 );
5179 let g = vec![0.01f32; inter * hidden];
5180 let d = vec![0.01f32; hidden * inter];
5181 add_raw(
5182 &format!("{p}.layers.0.mlp.gate_proj.weight"),
5183 vec![inter, hidden],
5184 &g,
5185 );
5186 add_raw(
5187 &format!("{p}.layers.0.mlp.up_proj.weight"),
5188 vec![inter, hidden],
5189 &g,
5190 );
5191 add_raw(
5192 &format!("{p}.layers.0.mlp.down_proj.weight"),
5193 vec![hidden, inter],
5194 &d,
5195 );
5196 let packed = layers * ple_d;
5197 add_raw(
5198 &format!("{p}.embed_tokens_per_layer.weight"),
5199 vec![vocab, packed],
5200 &vec![0.02f32; vocab * packed],
5201 );
5202 add_raw(
5203 &format!("{p}.per_layer_model_projection.weight"),
5204 vec![packed, hidden],
5205 &vec![0.01f32; packed * hidden],
5206 );
5207 add_raw(
5208 &format!("{p}.per_layer_projection_norm.weight"),
5209 vec![ple_d],
5210 &vec![1.0f32; ple_d],
5211 );
5212 add_raw(
5213 &format!("{p}.layers.0.per_layer_input_gate.weight"),
5214 vec![ple_d, hidden],
5215 &vec![0.01f32; ple_d * hidden],
5216 );
5217 add_raw(
5218 &format!("{p}.layers.0.per_layer_projection.weight"),
5219 vec![hidden, ple_d],
5220 &vec![0.01f32; hidden * ple_d],
5221 );
5222 add_raw(
5223 &format!("{p}.layers.0.post_per_layer_input_norm.weight"),
5224 vec![hidden],
5225 &n1,
5226 );
5227 add_raw(
5228 &format!("{p}.layers.0.altup.modality_router.weight"),
5229 vec![n_alt, hidden],
5230 &vec![0.01f32; n_alt * hidden],
5231 );
5232 add_raw(
5233 &format!("{p}.layers.0.altup.router_norm.weight"),
5234 vec![hidden],
5235 &n1,
5236 );
5237 add_raw(
5238 &format!("{p}.layers.0.altup.prediction_coefs.weight"),
5239 vec![n_alt * n_alt, n_alt],
5240 &vec![0.0f32; n_alt * n_alt * n_alt],
5241 );
5242 add_raw(
5243 &format!("{p}.layers.0.altup.correction_coefs.weight"),
5244 vec![n_alt, n_alt],
5245 &vec![0.0f32; n_alt * n_alt],
5246 );
5247 add_raw(
5248 &format!("{p}.layers.0.altup.correct_output_scale"),
5249 vec![hidden],
5250 &n1,
5251 );
5252 add_raw(
5253 &format!("{p}.layers.0.laurel.linear_left.weight"),
5254 vec![rank, hidden],
5255 &vec![0.01f32; rank * hidden],
5256 );
5257 add_raw(
5258 &format!("{p}.layers.0.laurel.linear_right.weight"),
5259 vec![hidden, rank],
5260 &vec![0.01f32; hidden * rank],
5261 );
5262 add_raw(
5263 &format!("{p}.layers.0.laurel.post_laurel_norm.weight"),
5264 vec![hidden],
5265 &n1,
5266 );
5267 let eye: Vec<f32> = (0..hidden * hidden)
5268 .map(|i| if i / hidden == i % hidden { 0.05 } else { 0.0 })
5269 .collect();
5270 for i in 0..3 {
5271 add_raw(
5272 &format!("{p}.altup_projections.{i}.weight"),
5273 vec![hidden, hidden],
5274 &eye,
5275 );
5276 add_raw(
5277 &format!("{p}.altup_unembed_projections.{i}.weight"),
5278 vec![hidden, hidden],
5279 &eye,
5280 );
5281 }
5282 add_raw(&format!("{p}.norm.weight"), vec![hidden], &n1);
5283 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
5284
5285 let cfg = json!({
5286 "format": "aria-quant-bundle",
5287 "format_version": 2,
5288 "quantization": "test",
5289 "group_size_default": 32,
5290 "hadamard_seed": 0,
5291 "model": {
5292 "hidden_size": hidden,
5293 "num_layers": layers,
5294 "num_attention_heads": n_heads,
5295 "num_kv_heads": n_kv,
5296 "intermediate_size": inter,
5297 "vocab_size": vocab,
5298 "context_length": 32,
5299 "rope_theta": 1000000.0,
5300 "hidden_act": "gelu_pytorch_tanh",
5301 "tie_word_embeddings": true,
5302 "head_dim": head_dim,
5303 "sliding_window": 512,
5304 "layer_types": ["full_attention"]
5305 },
5306 "tensors": tensors
5307 });
5308 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
5309 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
5310
5311 let mut s = SessionBuilder::new()
5312 .model(dir.path())
5313 .family("gemma/gemma-3n-e2b-it")
5314 .build()
5315 .unwrap();
5316 assert!(s.has_gemma3n_graph());
5317 assert!(s.weights.ple.is_some());
5318 assert!(s.weights.layers[0].altup.is_some());
5319 assert!(s.weights.layers[0].laurel.is_some());
5320 let prompt = vec![1u32, 2];
5321 let batched = s
5322 .generate(
5323 &prompt,
5324 &GenerateOpts {
5325 max_tokens: 3,
5326 temperature: 0.0,
5327 },
5328 )
5329 .unwrap();
5330 let step = s
5331 .generate(
5332 &prompt,
5333 &GenerateOpts {
5334 max_tokens: 3,
5335 temperature: 0.0,
5336 },
5337 )
5338 .unwrap();
5339 assert_eq!(batched.tokens, step.tokens);
5340 assert_eq!(batched.tokens.len(), 3);
5341 }
5342
5343 #[test]
5344 fn gaussian_topk_sparsity_zeros_below_cutoff() {
5345 let row: Vec<f32> = (0..100).map(|i| i as f32).collect();
5346 let y = gaussian_topk(&row, 100, 0.95).unwrap();
5347 assert!(y.iter().all(|&v| v >= 0.0));
5348 assert!(y.iter().filter(|&&v| v == 0.0).count() > 50);
5349 assert!(y.iter().any(|&v| v > 0.0));
5350 }
5351
5352 #[test]
5353 fn gemma4_e2b_scale_missing_ple_is_hard_error() {
5354 let dir = tempfile::tempdir().unwrap();
5355 let hidden = 1024usize;
5356 let layers = 1usize;
5357 let vocab = 8usize;
5358 let n_heads = 8usize;
5359 let n_kv = 1usize;
5360 let head_dim = 128usize;
5361 let q_dim = n_heads * head_dim;
5362 let k_dim = n_kv * head_dim;
5363 let inter = 32usize;
5364 let p = "model.language_model";
5365
5366 let mut tensors = serde_json::Map::new();
5367 let mut bin = Vec::new();
5368 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
5369 let offset = bin.len();
5370 for &v in data {
5371 bin.extend_from_slice(&v.to_le_bytes());
5372 }
5373 let nbytes = data.len() * 4;
5374 tensors.insert(
5375 name.to_string(),
5376 json!({
5377 "kind": "raw",
5378 "dtype": "f32",
5379 "shape": shape,
5380 "offsets": { "data": [offset, nbytes] }
5381 }),
5382 );
5383 };
5384 let emb = vec![0.01f32; vocab * hidden];
5385 add_raw(
5386 &format!("{p}.embed_tokens.weight"),
5387 vec![vocab, hidden],
5388 &emb,
5389 );
5390 let ones = vec![1.0f32; hidden];
5391 add_raw(
5392 &format!("{p}.layers.0.input_layernorm.weight"),
5393 vec![hidden],
5394 &ones,
5395 );
5396 add_raw(
5397 &format!("{p}.layers.0.pre_feedforward_layernorm.weight"),
5398 vec![hidden],
5399 &ones,
5400 );
5401 add_raw(
5402 &format!("{p}.layers.0.post_attention_layernorm.weight"),
5403 vec![hidden],
5404 &ones,
5405 );
5406 add_raw(
5407 &format!("{p}.layers.0.post_feedforward_layernorm.weight"),
5408 vec![hidden],
5409 &ones,
5410 );
5411 add_raw(&format!("{p}.norm.weight"), vec![hidden], &ones);
5412 let q = vec![0.01f32; q_dim * hidden];
5413 let k = vec![0.01f32; k_dim * hidden];
5414 add_raw(
5415 &format!("{p}.layers.0.self_attn.q_proj.weight"),
5416 vec![q_dim, hidden],
5417 &q,
5418 );
5419 add_raw(
5420 &format!("{p}.layers.0.self_attn.k_proj.weight"),
5421 vec![k_dim, hidden],
5422 &k,
5423 );
5424 add_raw(
5425 &format!("{p}.layers.0.self_attn.v_proj.weight"),
5426 vec![k_dim, hidden],
5427 &k,
5428 );
5429 add_raw(
5430 &format!("{p}.layers.0.self_attn.o_proj.weight"),
5431 vec![hidden, q_dim],
5432 &q,
5433 );
5434 let qn = vec![1.0f32; head_dim];
5435 add_raw(
5436 &format!("{p}.layers.0.self_attn.q_norm.weight"),
5437 vec![head_dim],
5438 &qn,
5439 );
5440 add_raw(
5441 &format!("{p}.layers.0.self_attn.k_norm.weight"),
5442 vec![head_dim],
5443 &qn,
5444 );
5445 let g = vec![0.01f32; inter * hidden];
5446 add_raw(
5447 &format!("{p}.layers.0.mlp.gate_proj.weight"),
5448 vec![inter, hidden],
5449 &g,
5450 );
5451 add_raw(
5452 &format!("{p}.layers.0.mlp.up_proj.weight"),
5453 vec![inter, hidden],
5454 &g,
5455 );
5456 add_raw(
5457 &format!("{p}.layers.0.mlp.down_proj.weight"),
5458 vec![hidden, inter],
5459 &g,
5460 );
5461 let cfg = json!({
5462 "format": "aria-quant-bundle",
5463 "format_version": 2,
5464 "quantization": "test",
5465 "group_size_default": 32,
5466 "hadamard_seed": 0,
5467 "model": {
5468 "hidden_size": hidden,
5469 "num_layers": layers,
5470 "num_attention_heads": n_heads,
5471 "num_kv_heads": n_kv,
5472 "intermediate_size": inter,
5473 "vocab_size": vocab,
5474 "context_length": 32,
5475 "rope_theta": 10000.0,
5476 "hidden_act": "gelu_pytorch_tanh",
5477 "tie_word_embeddings": true,
5478 "head_dim": head_dim,
5479 "global_head_dim": head_dim,
5480 "sliding_window": 512,
5481 "partial_rotary_factor": 0.25,
5482 "num_kv_shared_layers": 0,
5483 "layer_types": ["full_attention"]
5484 },
5485 "tensors": tensors
5486 });
5487 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
5488 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
5489
5490 let err = SessionBuilder::new()
5491 .model(dir.path())
5492 .family("gemma/gemma-4-e2b-it")
5493 .build()
5494 .unwrap_err();
5495 let msg = err.to_string();
5496 assert!(
5497 msg.contains("PLE") && msg.contains("embed_tokens_per_layer"),
5498 "{msg}"
5499 );
5500 }
5501
5502 #[test]
5503 fn gemma4_sliding_window_config_and_generate() {
5504 let dir_wide = tempfile::tempdir().unwrap();
5505 write_tiny_q4_bundle(dir_wide.path()).unwrap();
5506 let dir_narrow = tempfile::tempdir().unwrap();
5507 write_tiny_q4_bundle(dir_narrow.path()).unwrap();
5508 let patch = |path: &std::path::Path, window: usize| {
5509 let cfg_path = path.join("config.json");
5510 let raw = std::fs::read_to_string(&cfg_path).unwrap();
5511 let mut cfg: Value = serde_json::from_str(&raw).unwrap();
5512 cfg["model"]["sliding_window"] = json!(window);
5513 cfg["model"]["layer_types"] = json!(["sliding_attention", "sliding_attention"]);
5514 std::fs::write(&cfg_path, cfg.to_string()).unwrap();
5515 };
5516 patch(dir_wide.path(), 512);
5517 patch(dir_narrow.path(), 1);
5518 let wide = SessionBuilder::new()
5519 .model(dir_wide.path())
5520 .family("gemma/gemma-4-e2b-it")
5521 .build()
5522 .unwrap();
5523 let mut narrow = SessionBuilder::new()
5524 .model(dir_narrow.path())
5525 .family("gemma/gemma-4-e2b-it")
5526 .build()
5527 .unwrap();
5528 assert_eq!(wide.config().sliding_window, Some(512));
5529 assert_eq!(narrow.config().sliding_window, Some(1));
5530 assert_eq!(wide.attn_window(AttnKind::Sliding), Some(512));
5531 assert_eq!(narrow.attn_window(AttnKind::Sliding), Some(1));
5532 for layer in &narrow.weights.layers {
5533 if let LayerOp::Attn(attn) = &layer.op {
5534 assert_eq!(attn.kind, AttnKind::Sliding);
5535 }
5536 }
5537 let prompt = vec![1u32, 2, 3, 4];
5538 let gen = narrow
5539 .generate(
5540 &prompt,
5541 &GenerateOpts {
5542 max_tokens: 3,
5543 temperature: 0.0,
5544 },
5545 )
5546 .unwrap();
5547 assert_eq!(gen.tokens.len(), 3);
5548
5549 let mut incr = SessionBuilder::new()
5551 .model(dir_narrow.path())
5552 .family("gemma/gemma-4-e2b-it")
5553 .build()
5554 .unwrap();
5555 let again = incr
5556 .generate(
5557 &prompt,
5558 &GenerateOpts {
5559 max_tokens: 3,
5560 temperature: 0.0,
5561 },
5562 )
5563 .unwrap();
5564 assert_eq!(gen.tokens, again.tokens);
5565 }
5566}