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::tensor_names::{
6 action_head_names, attn_k_names, attn_k_norm_names, attn_norm_names, attn_o_names,
7 attn_post_norm_names, attn_q_names, attn_q_norm_names, attn_v_names, attn_v_norm_names,
8 conv_in_proj_names, conv_kernel_names, conv_out_proj_names, emb_names, embed_per_layer_names,
9 ffn_down_names, ffn_gate_names, ffn_norm_names, ffn_post_norm_names, ffn_up_names,
10 layer_ple_gate_names, layer_ple_post_norm_names, layer_ple_proj_names, layer_scalar_names,
11 linear_a_log_names,
12 linear_conv1d_names, linear_dt_bias_names, linear_in_proj_a_names, linear_in_proj_b_names,
13 linear_in_proj_ba_names, linear_in_proj_qkv_names, linear_in_proj_qkvz_names,
14 linear_in_proj_z_names, linear_out_norm_names, linear_out_proj_names, moe_expert_down_names, moe_expert_gate_names, moe_expert_up_names,
15 moe_router_names, output_names, output_norm_names, per_layer_model_projection_names,
16 per_layer_projection_norm_names, pre_feedforward_norm_names, vision_proj_names,
17};
18use crate::profile::{
19 elapsed_ms, load_profile_begin, load_profile_set_cuda_upload, load_profile_set_materialize,
20 load_profile_set_mmap, load_profile_take, EngineProfile, GenerateProfile,
21};
22use crate::tokenizer::{decode_placeholders, encode_naive, BundleTokenizer};
23use aria_kernel::{
24 attention_causal_with_scale, attention_with_scale, gated_delta_step, geglu, gelu_pytorch_tanh,
25 hdm_linear, kv_sliding_view, linear_cpu, moe_topk_route, resolve_compute, rms_norm,
26 rms_norm_gemma, rope_half, rope_half_partial, rope_half_proportional, short_conv_step, silu_vec,
27 softplus, swiglu,
28 ComputeBackend, ComputePref, CudaContext, EngineError, GatedDeltaStep,
29};
30use std::cell::RefCell;
31use std::collections::HashMap;
32use std::path::Path;
33use std::sync::Arc;
34use std::time::Instant;
35
36#[derive(Debug, Clone)]
37pub struct GenerateOpts {
38 pub max_tokens: usize,
39 pub temperature: f32,
40}
41
42impl Default for GenerateOpts {
43 fn default() -> Self {
44 Self {
45 max_tokens: 16,
46 temperature: 0.0,
47 }
48 }
49}
50
51#[derive(Debug, Clone)]
52pub struct Generation {
53 pub tokens: Vec<u32>,
54 pub text: String,
55}
56
57#[derive(Clone)]
58struct MatWeight {
59 data: Arc<Vec<f32>>,
60 hdm_seed: Option<i64>,
61}
62
63impl MatWeight {
64 fn from_loaded(w: LoadedWeight) -> Self {
65 Self {
66 data: Arc::new(w.data),
67 hdm_seed: w.hdm_seed,
68 }
69 }
70
71 fn concat_out(a: &Self, b: &Self) -> Self {
73 let mut data = Vec::with_capacity(a.data.len() + b.data.len());
74 data.extend_from_slice(&a.data);
75 data.extend_from_slice(&b.data);
76 Self {
77 data: Arc::new(data),
78 hdm_seed: None,
79 }
80 }
81}
82
83#[derive(Clone, Copy)]
84enum GemmAcct {
85 Attn,
86 Ffn,
87 LmHead,
88 Other,
89}
90
91#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
92enum AttnKind {
93 Sliding,
94 Full,
95}
96
97#[derive(Debug, Clone, Copy, PartialEq)]
100enum RopeMode {
101 Full,
102 Partial(f32),
103 Proportional(f32),
104}
105
106struct DecodeState {
108 k_caches: Vec<Vec<f32>>,
109 v_caches: Vec<Vec<f32>>,
110 last_kv_src: HashMap<AttnKind, usize>,
111 conv_states: Vec<Option<Vec<f32>>>,
112 delta_states: Vec<Option<Vec<f32>>>,
113 pos: usize,
115}
116
117#[derive(Clone)]
118struct AttnWeights {
119 wq: MatWeight,
120 wk: Option<MatWeight>,
122 wv: Option<MatWeight>,
123 wo: MatWeight,
124 q_norm: Option<Vec<f32>>,
125 k_norm: Option<Vec<f32>>,
126 v_norm: Option<Vec<f32>>,
127 kind: AttnKind,
128 q_gate: bool,
131}
132
133#[derive(Clone)]
134struct ConvWeights {
135 in_proj: MatWeight,
136 out_proj: MatWeight,
137 kernel: Vec<f32>,
139 kernel_size: usize,
140}
141
142#[derive(Clone)]
144struct DeltaWeights {
145 qkvz: MatWeight,
146 ba: MatWeight,
147 conv: Vec<f32>,
148 conv_k: usize,
149 out_proj: MatWeight,
150 out_norm: Vec<f32>,
152 a_log: Vec<f32>,
153 dt_bias: Vec<f32>,
154 n_k_heads: usize,
155 n_v_heads: usize,
156 head_k: usize,
157 head_v: usize,
158}
159
160#[derive(Clone)]
161enum LayerOp {
162 Attn(AttnWeights),
163 Conv(ConvWeights),
164 Linear(DeltaWeights),
165}
166
167#[derive(Clone)]
168struct ExpertWeights {
169 gate: MatWeight,
170 up: MatWeight,
171 down: MatWeight,
172}
173
174#[derive(Clone)]
175enum FfnWeights {
176 Dense {
177 gate: MatWeight,
178 up: MatWeight,
179 down: MatWeight,
180 },
181 MoE {
182 router: MatWeight,
183 experts: Vec<ExpertWeights>,
184 top_k: usize,
185 use_sigmoid: bool,
186 },
187}
188
189struct LayerPle {
190 gate: MatWeight,
191 proj: MatWeight,
192 post_norm: Vec<f32>,
193}
194
195struct PleModel {
196 embed: Arc<Vec<f32>>,
197 proj: MatWeight,
198 proj_norm: Vec<f32>,
199 d: usize,
200}
201
202struct LayerWeights {
203 attn_norm: Vec<f32>,
204 ffn_norm: Vec<f32>,
205 post_attn_norm: Option<Vec<f32>>,
206 post_ffn_norm: Option<Vec<f32>>,
207 ple: Option<LayerPle>,
208 layer_scalar: f32,
210 op: LayerOp,
211 ffn: FfnWeights,
212}
213
214struct ModelWeights {
215 emb: MatWeight,
216 layers: Vec<LayerWeights>,
217 output_norm: Vec<f32>,
218 output: MatWeight,
219 vision: Option<MatWeight>,
220 action: Option<MatWeight>,
221 ple: Option<PleModel>,
222}
223
224pub struct Session {
225 family: Family,
226 bundle: Bundle,
227 weights: ModelWeights,
228 conf: crate::bundle::ModelConfig,
229 use_gemma_norm: bool,
230 use_gemma4: bool,
231 use_geglu: bool,
232 embed_scale: f32,
233 final_logit_softcap: Option<f32>,
235 tokenizer: Option<BundleTokenizer>,
236 decode: Option<DecodeState>,
238 compute: ComputeBackend,
239 compute_label: String,
240 cuda: Option<CudaContext>,
241 profile_on: bool,
242 last_profile: Option<EngineProfile>,
243 gen_acc: RefCell<GenerateProfile>,
244}
245
246impl std::fmt::Debug for Session {
247 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
248 f.debug_struct("Session")
249 .field("family", &self.family)
250 .field("model", &self.conf.hidden_size)
251 .finish()
252 }
253}
254
255pub struct SessionBuilder {
256 path: Option<std::path::PathBuf>,
257 family_path: String,
258 compute: ComputePref,
259 profile: bool,
260}
261
262impl SessionBuilder {
263 pub fn new() -> Self {
264 Self {
265 path: None,
266 family_path: "gemma/gemma-4-e2b-it".into(),
267 compute: ComputePref::Auto,
268 profile: false,
269 }
270 }
271
272 pub fn model(mut self, path: impl AsRef<Path>) -> Self {
273 self.path = Some(path.as_ref().to_path_buf());
274 self
275 }
276
277 pub fn family(mut self, path: impl Into<String>) -> Self {
278 self.family_path = path.into();
279 self
280 }
281
282 pub fn compute(mut self, pref: ComputePref) -> Self {
283 self.compute = pref;
284 self
285 }
286
287 pub fn profile(mut self, on: bool) -> Self {
288 self.profile = on;
289 self
290 }
291
292 pub fn build(self) -> Result<Session, EngineError> {
293 let family = require_runnable(&self.family_path)?;
294 let _hook = graph_hook(family.arch);
295 let path = self
296 .path
297 .ok_or_else(|| EngineError::InvalidParam("model path required".into()))?;
298 let (compute, compute_label) = resolve_compute(self.compute)?;
299 load_profile_begin(self.profile);
300 let t_mmap = Instant::now();
301 let bundle = load_bundle(&path)?;
302 load_profile_set_mmap(elapsed_ms(t_mmap));
303 let mut conf = bundle.model.clone();
304 conf.rope_theta = effective_rope_theta(family.path(), conf.rope_theta);
305 fill_gemma4_architecture_defaults(&mut conf, family.path());
308 fill_qwen35_architecture_defaults(&mut conf, family.path());
309 require_gemma4_config(&conf, family.path())?;
310 reject_unsupported_geometry(&conf, family)?;
311 let tokenizer = BundleTokenizer::try_load(&path)?;
312 let t_mat = Instant::now();
313 let weights = materialize_with_config(&bundle, family, &conf)?;
316 require_gemma4_ple(&weights, &conf, family.path())?;
317 load_profile_set_materialize(elapsed_ms(t_mat));
318 let mut cuda = None;
319 if compute == ComputeBackend::Cuda {
320 let t_up = Instant::now();
321 let ctx = CudaContext::new()?;
322 upload_weights(&ctx, &weights)?;
323 load_profile_set_cuda_upload(elapsed_ms(t_up));
324 cuda = Some(ctx);
325 }
326 let act = conf
327 .hidden_act
328 .as_deref()
329 .unwrap_or("")
330 .to_ascii_lowercase();
331 let use_gemma4 = family.path().contains("gemma-4");
332 let use_gemma_norm = (family.path().contains("gemma") && !use_gemma4)
334 || family.path().contains("qwen3.5");
335 let use_geglu = act.contains("gelu") || use_gemma4;
336 let embed_scale = if family.path().contains("gemma") {
339 (conf.hidden_size as f32).sqrt()
340 } else {
341 1.0
342 };
343 let final_logit_softcap = if use_gemma4 { Some(30.0) } else { None };
344 let load = load_profile_take();
345 let last_profile = self.profile.then(|| EngineProfile {
346 compute: compute_label.clone(),
347 load,
348 generate: None,
349 ci_fail: false,
350 });
351 Ok(Session {
352 family,
353 bundle,
354 weights,
355 conf,
356 use_gemma_norm,
357 use_gemma4,
358 use_geglu,
359 embed_scale,
360 final_logit_softcap,
361 tokenizer,
362 decode: None,
363 compute,
364 compute_label,
365 cuda,
366 profile_on: self.profile,
367 last_profile,
368 gen_acc: RefCell::new(GenerateProfile::default()),
369 })
370 }
371}
372
373impl Default for SessionBuilder {
374 fn default() -> Self {
375 Self::new()
376 }
377}
378
379fn reject_unsupported_geometry(
380 conf: &crate::bundle::ModelConfig,
381 family: Family,
382) -> Result<(), EngineError> {
383 let path = family.path();
384 if path.contains("qwen3.5") || path.contains("bonsai") {
386 let has_linear = conf
387 .layer_types
388 .as_ref()
389 .map(|t| {
390 t.iter().any(|s| {
391 let s = s.to_ascii_lowercase();
392 s.contains("linear_attention") || s.contains("delta")
393 })
394 })
395 .unwrap_or(false);
396 if !has_linear {
397 return Err(EngineError::Unsupported(format!(
398 "{path}: requires model.layer_types with Gated DeltaNet / linear_attention \
399 (dense-only bundles are unsupported until DeltaNet lands)"
400 )));
401 }
402 }
403 if family.is_moe() && conf.num_experts.unwrap_or(0) == 0 {
405 return Err(EngineError::Unsupported(format!(
406 "{path}: MoE family requires model.num_experts > 0 in bundle config"
407 )));
408 }
409 Ok(())
410}
411
412fn layer_type_str(conf: &crate::bundle::ModelConfig, layer: usize) -> String {
413 conf.layer_types
414 .as_ref()
415 .and_then(|t| t.get(layer))
416 .map(|s| s.to_ascii_lowercase())
417 .unwrap_or_else(|| "full_attention".into())
418}
419
420fn layer_is_conv(conf: &crate::bundle::ModelConfig, layer: usize) -> bool {
421 layer_type_str(conf, layer).contains("conv")
422}
423
424fn layer_is_linear(conf: &crate::bundle::ModelConfig, layer: usize) -> bool {
425 let t = layer_type_str(conf, layer);
426 t.contains("linear_attention") || t.contains("delta")
427}
428
429fn attn_kind(conf: &crate::bundle::ModelConfig, layer: usize) -> AttnKind {
430 if layer_type_str(conf, layer).contains("sliding") {
431 AttnKind::Sliding
432 } else {
433 AttnKind::Full
434 }
435}
436
437fn is_kv_consumer(conf: &crate::bundle::ModelConfig, layer: usize) -> bool {
438 let n = conf.num_kv_shared_layers.unwrap_or(0);
439 n > 0 && layer >= conf.num_layers.saturating_sub(n)
440}
441
442fn default_gemma4_layer_types(n: usize) -> Vec<String> {
444 (0..n)
445 .map(|i| {
446 if (i + 1) % 5 == 0 {
447 "full_attention".into()
448 } else {
449 "sliding_attention".into()
450 }
451 })
452 .collect()
453}
454
455fn fill_qwen35_architecture_defaults(conf: &mut crate::bundle::ModelConfig, family_path: &str) {
459 if !family_path.contains("qwen3.5") {
460 return;
461 }
462 if conf.partial_rotary_factor.is_none() {
463 conf.partial_rotary_factor = Some(0.25);
464 }
465}
466
467fn fill_gemma4_architecture_defaults(conf: &mut crate::bundle::ModelConfig, family_path: &str) {
470 if !family_path.contains("gemma-4") {
471 return;
472 }
473 if conf.layer_types.as_ref().map(|t| t.len()) != Some(conf.num_layers) {
474 conf.layer_types = Some(default_gemma4_layer_types(conf.num_layers));
475 }
476 if conf.sliding_window.unwrap_or(0) == 0 {
477 conf.sliding_window = Some(512);
478 }
479 if conf.partial_rotary_factor.is_none() {
480 conf.partial_rotary_factor = Some(0.25);
481 }
482 if conf.hidden_size >= 1024 {
484 if conf.head_dim.unwrap_or(0) == 0 {
485 conf.head_dim = Some(256);
486 }
487 if conf.global_head_dim.unwrap_or(0) == 0 {
488 conf.global_head_dim = Some(512);
489 }
490 if conf.num_kv_shared_layers.is_none() {
491 conf.num_kv_shared_layers = Some(20);
492 }
493 } else {
494 if conf.head_dim.unwrap_or(0) == 0 && conf.num_attention_heads > 0 {
495 conf.head_dim = Some(conf.hidden_size / conf.num_attention_heads);
496 }
497 if conf.global_head_dim.unwrap_or(0) == 0 {
498 conf.global_head_dim = conf.head_dim;
499 }
500 }
501}
502
503fn require_gemma4_config(
505 conf: &crate::bundle::ModelConfig,
506 family_path: &str,
507) -> Result<(), EngineError> {
508 if !family_path.contains("gemma-4") {
509 return Ok(());
510 }
511 let missing = |field: &str| {
512 EngineError::Unsupported(format!(
513 "{family_path}: model.{field} required after Gemma-4 architecture fill \
514 (re-quantize with current model config_from_hf)"
515 ))
516 };
517 match &conf.layer_types {
518 None => return Err(missing("layer_types")),
519 Some(t) if t.len() != conf.num_layers => {
520 return Err(EngineError::Unsupported(format!(
521 "{family_path}: model.layer_types length {} != num_layers {}",
522 t.len(),
523 conf.num_layers
524 )));
525 }
526 Some(_) => {}
527 }
528 if conf.sliding_window.unwrap_or(0) == 0 {
529 return Err(missing("sliding_window"));
530 }
531 match conf.partial_rotary_factor {
532 Some(f) if f > 0.0 && f <= 1.0 => {}
533 _ => return Err(missing("partial_rotary_factor")),
534 }
535 if conf.head_dim.unwrap_or(0) == 0 {
536 return Err(missing("head_dim"));
537 }
538 if conf.global_head_dim.unwrap_or(0) == 0 {
539 return Err(missing("global_head_dim"));
540 }
541 Ok(())
542}
543
544fn gemma4_requires_ple(family_path: &str, hidden: usize) -> bool {
545 family_path.contains("gemma-4") && hidden >= 1024
546}
547
548fn require_gemma4_ple(
549 weights: &ModelWeights,
550 conf: &crate::bundle::ModelConfig,
551 family_path: &str,
552) -> Result<(), EngineError> {
553 if !gemma4_requires_ple(family_path, conf.hidden_size) {
554 return Ok(());
555 }
556 if weights.ple.is_none() {
557 return Err(EngineError::Format(format!(
558 "{family_path}: codebook PLE required for Gemma-4 E2B/E4B \
559 (embed_tokens_per_layer + per_layer_model_projection + \
560 per_layer_projection_norm); refusing silent no-op"
561 )));
562 }
563 Ok(())
564}
565
566fn resolve_attn_kind(
568 conf: &crate::bundle::ModelConfig,
569 layer: usize,
570 q_dim: usize,
571 n_heads: usize,
572) -> AttnKind {
573 if n_heads > 0 && q_dim.is_multiple_of(n_heads) {
574 let head_from_q = q_dim / n_heads;
575 if let (Some(g), Some(h)) = (
576 conf.global_head_dim.filter(|d| *d > 0),
577 conf.head_dim.filter(|d| *d > 0),
578 ) {
579 if g != h {
580 if head_from_q == g {
581 return AttnKind::Full;
582 }
583 if head_from_q == h {
584 return AttnKind::Sliding;
585 }
586 }
587 }
588 }
589 attn_kind(conf, layer)
590}
591
592fn attn_q_geometry(
594 wq_len: usize,
595 wo_len: usize,
596 hidden: usize,
597) -> Result<(usize, usize, bool), EngineError> {
598 if hidden == 0 || !wq_len.is_multiple_of(hidden) {
599 return Err(EngineError::ShapeMismatch(format!(
600 "attn q proj weight not divisible by hidden_size (len={wq_len} hidden={hidden})"
601 )));
602 }
603 if !wo_len.is_multiple_of(hidden) {
604 return Err(EngineError::ShapeMismatch(format!(
605 "attn output proj weight not divisible by hidden_size (len={wo_len} hidden={hidden})"
606 )));
607 }
608 let q_out = wq_len / hidden;
609 let wo_in = wo_len / hidden;
610 if q_out == wo_in {
611 Ok((q_out, q_out, false))
612 } else if q_out == 2 * wo_in {
613 Ok((q_out, wo_in, true))
614 } else {
615 Err(EngineError::ShapeMismatch(format!(
616 "attn output proj weight shape mismatch (wo_len={wo_len} hidden={hidden} q_out={q_out})"
617 )))
618 }
619}
620
621fn split_interleaved_q_gate(
623 mixed: &[f32],
624 seq: usize,
625 n_heads: usize,
626 head_dim: usize,
627) -> Result<(Vec<f32>, Vec<f32>), EngineError> {
628 let q_dim = n_heads.saturating_mul(head_dim);
629 let packed = q_dim.saturating_mul(2);
630 if packed == 0 || mixed.len() != seq * packed {
631 return Err(EngineError::ShapeMismatch(format!(
632 "gated q_proj out {} != seq*2*q_dim {}*{packed}",
633 mixed.len(),
634 seq
635 )));
636 }
637 let mut q = vec![0.0f32; seq * q_dim];
638 let mut gate = vec![0.0f32; seq * q_dim];
639 for t in 0..seq {
640 for h in 0..n_heads {
641 let src = t * packed + h * (2 * head_dim);
642 let dst = t * q_dim + h * head_dim;
643 q[dst..dst + head_dim].copy_from_slice(&mixed[src..src + head_dim]);
644 gate[dst..dst + head_dim]
645 .copy_from_slice(&mixed[src + head_dim..src + 2 * head_dim]);
646 }
647 }
648 Ok((q, gate))
649}
650
651fn apply_sigmoid_gate(x: &mut [f32], gate: &[f32]) -> Result<(), EngineError> {
652 if x.len() != gate.len() {
653 return Err(EngineError::ShapeMismatch(format!(
654 "attn output gate len {} != attn out {}",
655 gate.len(),
656 x.len()
657 )));
658 }
659 for (v, g) in x.iter_mut().zip(gate) {
660 *v *= 1.0 / (1.0 + (-*g).exp());
661 }
662 Ok(())
663}
664
665fn materialize_with_config(
666 b: &Bundle,
667 family: Family,
668 conf: &crate::bundle::ModelConfig,
669) -> Result<ModelWeights, EngineError> {
670 fn any_mat(b: &Bundle, names: &[String]) -> Result<MatWeight, EngineError> {
671 let refs: Vec<&str> = names.iter().map(String::as_str).collect();
672 Ok(MatWeight::from_loaded(b.weight_loaded_any(&refs)?))
673 }
674 fn try_mat(b: &Bundle, names: &[String]) -> Result<Option<MatWeight>, EngineError> {
675 match any_mat(b, names) {
676 Ok(w) => Ok(Some(w)),
677 Err(EngineError::Format(_)) => Ok(None),
678 Err(e) => Err(e),
679 }
680 }
681 fn any_vec(b: &Bundle, names: &[String]) -> Result<Vec<f32>, EngineError> {
682 Ok((*any_mat(b, names)?.data).clone())
683 }
684 fn optional_vec(b: &Bundle, names: &[String]) -> Option<Vec<f32>> {
685 any_vec(b, names).ok()
686 }
687
688 let m = conf;
689 let hidden = m.hidden_size;
690 let n_heads = m.num_attention_heads;
691 let n_experts = m.num_experts.unwrap_or(0);
692 let top_k = m.num_experts_per_tok.unwrap_or(1).max(1);
693 let use_sigmoid_router = n_experts > 0 && m.layer_types.is_some();
695
696 let mut layers = Vec::with_capacity(m.num_layers);
697 let mut prev_wk: Option<MatWeight> = None;
698 let mut prev_wv: Option<MatWeight> = None;
699 for layer in 0..m.num_layers {
700 let attn_norm = any_vec(b, &attn_norm_names(layer))?;
701 let pre_ff = optional_vec(b, &pre_feedforward_norm_names(layer));
702 let post_attn_norm = if pre_ff.is_some() {
703 optional_vec(b, &attn_post_norm_names(layer))
704 } else {
705 None
706 };
707 let post_ffn_norm = optional_vec(b, &ffn_post_norm_names(layer));
708 let ffn_norm = if let Some(v) = pre_ff {
709 v
710 } else {
711 any_vec(b, &ffn_norm_names(layer))?
712 };
713
714 let op = if layer_is_conv(m, layer) {
715 let in_proj = any_mat(b, &conv_in_proj_names(layer))?;
716 let out_proj = any_mat(b, &conv_out_proj_names(layer))?;
717 let kw = any_mat(b, &conv_kernel_names(layer))?;
718 let kernel_size = m.conv_l_cache.unwrap_or(3).max(1);
719 if kw.data.len() % hidden != 0 {
720 return Err(EngineError::ShapeMismatch(format!(
721 "layer {layer} conv kernel len {} not divisible by hidden {hidden}",
722 kw.data.len()
723 )));
724 }
725 let inferred_k = kw.data.len() / hidden;
726 let kernel_size = if inferred_k > 0 {
727 inferred_k
728 } else {
729 kernel_size
730 };
731 if kw.data.len() != hidden * kernel_size {
733 return Err(EngineError::ShapeMismatch(format!(
734 "layer {layer} conv kernel len {} != hidden*kernel {hidden}*{kernel_size}",
735 kw.data.len()
736 )));
737 }
738 let kernel = (*kw.data).clone();
739 if in_proj.data.len() != 3 * hidden * hidden {
740 return Err(EngineError::ShapeMismatch(format!(
741 "layer {layer} conv in_proj len {} != 3*hidden*hidden",
742 in_proj.data.len()
743 )));
744 }
745 if out_proj.data.len() != hidden * hidden {
746 return Err(EngineError::ShapeMismatch(format!(
747 "layer {layer} conv out_proj len {} != hidden*hidden",
748 out_proj.data.len()
749 )));
750 }
751 LayerOp::Conv(ConvWeights {
752 in_proj,
753 out_proj,
754 kernel,
755 kernel_size,
756 })
757 } else if layer_is_linear(m, layer) {
758 let qkvz = if let Some(w) = try_mat(b, &linear_in_proj_qkvz_names(layer))? {
759 w
760 } else {
761 let qkv = any_mat(b, &linear_in_proj_qkv_names(layer))?;
762 let z = any_mat(b, &linear_in_proj_z_names(layer))?;
763 MatWeight::concat_out(&qkv, &z)
764 };
765 let ba = if let Some(w) = try_mat(b, &linear_in_proj_ba_names(layer))? {
766 w
767 } else {
768 let proj_b = any_mat(b, &linear_in_proj_b_names(layer))?;
769 let proj_a = any_mat(b, &linear_in_proj_a_names(layer))?;
770 MatWeight::concat_out(&proj_b, &proj_a)
771 };
772 let conv_w = any_mat(b, &linear_conv1d_names(layer))?;
773 let out_proj = any_mat(b, &linear_out_proj_names(layer))?;
774 let a_log = any_vec(b, &linear_a_log_names(layer))?;
775 let dt_bias = any_vec(b, &linear_dt_bias_names(layer))?;
776 let n_v_heads = a_log.len();
777 if n_v_heads == 0 || dt_bias.len() != n_v_heads {
778 return Err(EngineError::ShapeMismatch(format!(
779 "layer {layer} A_log/dt_bias head mismatch"
780 )));
781 }
782 if ba.data.len() % hidden != 0 {
783 return Err(EngineError::ShapeMismatch(
784 "linear in_proj_ba not divisible by hidden".into(),
785 ));
786 }
787 if ba.data.len() / hidden != 2 * n_v_heads {
788 return Err(EngineError::ShapeMismatch(format!(
789 "layer {layer} in_proj_ba out {} != 2*n_v_heads {}",
790 ba.data.len() / hidden,
791 2 * n_v_heads
792 )));
793 }
794 if qkvz.data.len() % hidden != 0 {
795 return Err(EngineError::ShapeMismatch(
796 "linear in_proj_qkvz not divisible by hidden".into(),
797 ));
798 }
799 let qkvz_out = qkvz.data.len() / hidden;
800 if !qkvz_out.is_multiple_of(4) {
802 return Err(EngineError::ShapeMismatch(format!(
803 "layer {layer} qkvz out {qkvz_out} not divisible by 4"
804 )));
805 }
806 let key_dim = qkvz_out / 4;
807 let value_dim = key_dim;
808 let n_k_heads = n_v_heads;
809 if n_k_heads == 0
810 || !key_dim.is_multiple_of(n_k_heads)
811 || !value_dim.is_multiple_of(n_v_heads)
812 {
813 return Err(EngineError::ShapeMismatch(format!(
814 "layer {layer} cannot infer DeltaNet head dims"
815 )));
816 }
817 let head_k = key_dim / n_k_heads;
818 let head_v = value_dim / n_v_heads;
819 let conv_dim = key_dim * 2 + value_dim;
820 if conv_w.data.len() % conv_dim != 0 {
821 return Err(EngineError::ShapeMismatch(format!(
822 "layer {layer} conv1d len {} not divisible by conv_dim {conv_dim}",
823 conv_w.data.len()
824 )));
825 }
826 let conv_k = conv_w.data.len() / conv_dim;
827 if out_proj.data.len() != hidden * value_dim {
828 return Err(EngineError::ShapeMismatch(format!(
829 "layer {layer} linear out_proj len {} != hidden*value_dim",
830 out_proj.data.len()
831 )));
832 }
833 let out_norm = optional_vec(b, &linear_out_norm_names(layer))
834 .unwrap_or_else(|| vec![1.0f32; head_v]);
835 if out_norm.len() != head_v {
836 return Err(EngineError::ShapeMismatch(format!(
837 "layer {layer} linear_attn.norm len {} != head_v {head_v}",
838 out_norm.len()
839 )));
840 }
841 LayerOp::Linear(DeltaWeights {
842 qkvz,
843 ba,
844 conv: (*conv_w.data).clone(),
845 conv_k,
846 out_proj,
847 out_norm,
848 a_log,
849 dt_bias,
850 n_k_heads,
851 n_v_heads,
852 head_k,
853 head_v,
854 })
855 } else {
856 let consumer = is_kv_consumer(m, layer);
857 let (wk, wv) = if consumer {
858 (None, None)
859 } else {
860 let wk = match any_mat(b, &attn_k_names(layer)) {
861 Ok(w) => {
862 prev_wk = Some(w.clone());
863 Some(w)
864 }
865 Err(e) => Some(prev_wk.clone().ok_or_else(|| {
866 EngineError::Format(format!(
867 "missing k_proj for layer {layer} and no prior KV to share ({e})"
868 ))
869 })?),
870 };
871 let wv = match any_mat(b, &attn_v_names(layer)) {
872 Ok(w) => {
873 prev_wv = Some(w.clone());
874 Some(w)
875 }
876 Err(e) => Some(prev_wv.clone().ok_or_else(|| {
877 EngineError::Format(format!(
878 "missing v_proj for layer {layer} and no prior KV to share ({e})"
879 ))
880 })?),
881 };
882 (wk, wv)
883 };
884 let wq = any_mat(b, &attn_q_names(layer))?;
885 let wo = any_mat(b, &attn_o_names(layer))?;
886 let (_q_out, q_dim, q_gate) =
887 attn_q_geometry(wq.data.len(), wo.data.len(), hidden)?;
888 LayerOp::Attn(AttnWeights {
889 wq,
890 wk,
891 wv,
892 wo,
893 q_norm: optional_vec(b, &attn_q_norm_names(layer)),
894 k_norm: optional_vec(b, &attn_k_norm_names(layer)),
895 v_norm: optional_vec(b, &attn_v_norm_names(layer)),
896 kind: resolve_attn_kind(m, layer, q_dim, n_heads),
897 q_gate,
898 })
899 };
900
901 let ffn = if n_experts > 0 {
902 match any_mat(b, &moe_router_names(layer)) {
903 Ok(router) => {
904 if router.data.len() != n_experts * hidden {
905 return Err(EngineError::ShapeMismatch(format!(
906 "layer {layer} MoE router len {} != num_experts*hidden {n_experts}*{hidden}",
907 router.data.len()
908 )));
909 }
910 let mut experts = Vec::with_capacity(n_experts);
911 for e in 0..n_experts {
912 experts.push(ExpertWeights {
913 gate: any_mat(b, &moe_expert_gate_names(layer, e))?,
914 up: any_mat(b, &moe_expert_up_names(layer, e))?,
915 down: any_mat(b, &moe_expert_down_names(layer, e))?,
916 });
917 }
918 FfnWeights::MoE {
919 router,
920 experts,
921 top_k,
922 use_sigmoid: use_sigmoid_router,
923 }
924 }
925 Err(_) => {
926 FfnWeights::Dense {
928 gate: any_mat(b, &ffn_gate_names(layer))?,
929 up: any_mat(b, &ffn_up_names(layer))?,
930 down: any_mat(b, &ffn_down_names(layer))?,
931 }
932 }
933 }
934 } else {
935 FfnWeights::Dense {
936 gate: any_mat(b, &ffn_gate_names(layer))?,
937 up: any_mat(b, &ffn_up_names(layer))?,
938 down: any_mat(b, &ffn_down_names(layer))?,
939 }
940 };
941
942 let ple = match (
943 any_mat(b, &layer_ple_gate_names(layer)),
944 any_mat(b, &layer_ple_proj_names(layer)),
945 optional_vec(b, &layer_ple_post_norm_names(layer)),
946 ) {
947 (Ok(gate), Ok(proj), Some(post_norm)) => Some(LayerPle {
948 gate,
949 proj,
950 post_norm,
951 }),
952 _ => None,
953 };
954
955 let layer_scalar = optional_vec(b, &layer_scalar_names(layer))
956 .and_then(|v| v.into_iter().find(|x| x.is_finite()))
957 .unwrap_or(1.0);
958
959 layers.push(LayerWeights {
960 attn_norm,
961 ffn_norm,
962 post_attn_norm,
963 post_ffn_norm,
964 ple,
965 layer_scalar,
966 op,
967 ffn,
968 });
969 }
970 let emb_n = emb_names();
971 let out_norm_n = output_norm_names();
972 let out_n = output_names();
973 let vis_n: Vec<String> = vision_proj_names()
974 .iter()
975 .map(|s| (*s).to_string())
976 .collect();
977 let act_n: Vec<String> = action_head_names()
978 .iter()
979 .map(|s| (*s).to_string())
980 .collect();
981 let emb = MatWeight::from_loaded(b.weight_loaded_any(&emb_n)?);
982 let output = if m.tie_word_embeddings.unwrap_or(false)
983 || family.path().contains("gemma-4")
984 || (family.path().contains("qwen3") && !family.path().contains("qwen3.5"))
985 {
986 emb.clone()
989 } else {
990 MatWeight::from_loaded(b.weight_loaded_any(&out_n)?)
991 };
992 let require_ple = gemma4_requires_ple(family.path(), hidden);
993 let ple = {
994 let embed_n = embed_per_layer_names();
995 let proj_n: Vec<String> = per_layer_model_projection_names()
996 .iter()
997 .map(|s| (*s).to_string())
998 .collect();
999 let norm_n: Vec<String> = per_layer_projection_norm_names()
1000 .iter()
1001 .map(|s| (*s).to_string())
1002 .collect();
1003 let embed_res = b.weight_loaded_any(&embed_n);
1004 let proj_res = any_mat(b, &proj_n);
1005 let proj_norm = optional_vec(b, &norm_n);
1006 match (embed_res, proj_res, proj_norm) {
1007 (Ok(embed), Ok(proj), Some(proj_norm)) => {
1008 let d = proj_norm.len();
1009 if d == 0 {
1010 return Err(EngineError::ShapeMismatch(
1011 "PLE projection norm dim is 0".into(),
1012 ));
1013 }
1014 Some(PleModel {
1015 embed: Arc::new(embed.data),
1016 proj,
1017 proj_norm,
1018 d,
1019 })
1020 }
1021 (embed_res, proj_res, proj_norm) if require_ple => {
1022 let embed_s = match &embed_res {
1023 Ok(_) => "ok".to_string(),
1024 Err(e) => e.to_string(),
1025 };
1026 let proj_s = match &proj_res {
1027 Ok(_) => "ok".to_string(),
1028 Err(e) => e.to_string(),
1029 };
1030 let norm_s = if proj_norm.is_some() { "ok" } else { "missing" };
1031 return Err(EngineError::Format(format!(
1032 "{}: codebook PLE required (embed_tokens_per_layer={embed_s}, \
1033 per_layer_model_projection={proj_s}, per_layer_projection_norm={norm_s})",
1034 family.path()
1035 )));
1036 }
1037 _ => None,
1038 }
1039 };
1040 if let Some(ple) = &ple {
1041 let packed = m.num_layers.saturating_mul(ple.d);
1042 if packed == 0
1043 || !ple.embed.len().is_multiple_of(packed)
1044 || ple.proj.data.len() != packed * hidden
1045 {
1046 return Err(EngineError::ShapeMismatch(format!(
1047 "PLE shapes: embed {} proj {} expected packed={} hidden={hidden}",
1048 ple.embed.len(),
1049 ple.proj.data.len(),
1050 packed
1051 )));
1052 }
1053 for (i, layer) in layers.iter().enumerate() {
1054 let Some(lp) = &layer.ple else {
1055 return Err(EngineError::Format(format!(
1056 "PLE model tensors present but layer {i} missing gate/proj/norm"
1057 )));
1058 };
1059 if lp.gate.data.len() != ple.d * hidden || lp.proj.data.len() != hidden * ple.d {
1060 return Err(EngineError::ShapeMismatch(format!(
1061 "layer {i} PLE gate/proj shape mismatch (d={}, hidden={hidden})",
1062 ple.d
1063 )));
1064 }
1065 if lp.post_norm.len() != hidden {
1068 return Err(EngineError::ShapeMismatch(format!(
1069 "layer {i} PLE post_norm len {} != hidden {hidden}",
1070 lp.post_norm.len()
1071 )));
1072 }
1073 }
1074 }
1075 Ok(ModelWeights {
1076 emb,
1077 layers,
1078 output_norm: b.weight_loaded_any(&out_norm_n)?.data,
1079 output,
1080 vision: any_mat(b, &vis_n).ok(),
1081 action: any_mat(b, &act_n).ok(),
1082 ple,
1083 })
1084}
1085
1086fn upload_weights(ctx: &CudaContext, w: &ModelWeights) -> Result<(), EngineError> {
1087 ctx.upload(&w.emb.data)?;
1088 ctx.upload(&w.output.data)?;
1089 if let Some(v) = &w.vision {
1090 ctx.upload(&v.data)?;
1091 }
1092 if let Some(a) = &w.action {
1093 ctx.upload(&a.data)?;
1094 }
1095 if let Some(ple) = &w.ple {
1096 ctx.upload(&ple.embed)?;
1097 ctx.upload(&ple.proj.data)?;
1098 }
1099 for layer in &w.layers {
1100 match &layer.op {
1101 LayerOp::Attn(attn) => {
1102 ctx.upload(&attn.wq.data)?;
1103 if let Some(wk) = &attn.wk {
1104 ctx.upload(&wk.data)?;
1105 }
1106 if let Some(wv) = &attn.wv {
1107 ctx.upload(&wv.data)?;
1108 }
1109 ctx.upload(&attn.wo.data)?;
1110 }
1111 LayerOp::Conv(c) => {
1112 ctx.upload(&c.in_proj.data)?;
1113 ctx.upload(&c.out_proj.data)?;
1114 }
1115 LayerOp::Linear(d) => {
1116 ctx.upload(&d.qkvz.data)?;
1117 ctx.upload(&d.ba.data)?;
1118 ctx.upload(&d.out_proj.data)?;
1119 }
1120 }
1121 if let Some(ple) = &layer.ple {
1122 ctx.upload(&ple.gate.data)?;
1123 ctx.upload(&ple.proj.data)?;
1124 }
1125 match &layer.ffn {
1126 FfnWeights::Dense { gate, up, down } => {
1127 ctx.upload(&gate.data)?;
1128 ctx.upload(&up.data)?;
1129 ctx.upload(&down.data)?;
1130 }
1131 FfnWeights::MoE { router, experts, .. } => {
1132 ctx.upload(&router.data)?;
1133 for e in experts {
1134 ctx.upload(&e.gate.data)?;
1135 ctx.upload(&e.up.data)?;
1136 ctx.upload(&e.down.data)?;
1137 }
1138 }
1139 }
1140 }
1141 Ok(())
1142}
1143
1144impl Session {
1145 pub fn family(&self) -> Family {
1146 self.family
1147 }
1148
1149 pub fn model_id(&self) -> &str {
1150 self.family.path()
1151 }
1152
1153 pub fn config(&self) -> &crate::bundle::ModelConfig {
1154 &self.conf
1155 }
1156
1157 pub fn bundle(&self) -> &Bundle {
1158 &self.bundle
1159 }
1160
1161 pub fn compute_label(&self) -> &str {
1162 &self.compute_label
1163 }
1164
1165 pub fn last_profile(&self) -> Option<&EngineProfile> {
1166 self.last_profile.as_ref()
1167 }
1168
1169 fn wmm(
1170 &self,
1171 w: &MatWeight,
1172 x: &[f32],
1173 out_f: usize,
1174 in_f: usize,
1175 acct: GemmAcct,
1176 ) -> Result<Vec<f32>, EngineError> {
1177 let t0 = Instant::now();
1178 let y = if let Some(seed) = w.hdm_seed {
1179 hdm_linear(x, &w.data, out_f, in_f, Some(seed))?
1180 } else if self.compute == ComputeBackend::Cuda {
1181 let ctx = self.cuda.as_ref().ok_or_else(|| {
1182 EngineError::Unsupported("compute=cuda but CudaContext missing".into())
1183 })?;
1184 ctx.linear(x, &w.data, out_f, in_f)?
1185 } else {
1186 linear_cpu(x, &w.data, out_f, in_f)?
1187 };
1188 if self.profile_on {
1189 let ms = elapsed_ms(t0);
1190 let mut g = self.gen_acc.borrow_mut();
1191 match acct {
1192 GemmAcct::Attn => g.gemm_attn_ms += ms,
1193 GemmAcct::Ffn => g.gemm_ffn_ms += ms,
1194 GemmAcct::LmHead => g.gemm_lm_head_ms += ms,
1195 GemmAcct::Other => {}
1196 }
1197 }
1198 Ok(y)
1199 }
1200
1201 fn can_batch_prefill(&self) -> bool {
1202 self.weights.layers.iter().all(|layer| {
1203 matches!(layer.op, LayerOp::Attn(_)) && matches!(layer.ffn, FfnWeights::Dense { .. })
1204 })
1205 }
1206
1207 pub fn generate(
1210 &mut self,
1211 prompt: &[u32],
1212 opts: &GenerateOpts,
1213 ) -> Result<Generation, EngineError> {
1214 if opts.max_tokens == 0 {
1215 return Err(EngineError::InvalidParam("max_tokens must be > 0".into()));
1216 }
1217 let mut tokens: Vec<u32> = prompt.to_vec();
1218 if tokens.is_empty() {
1219 tokens.push(1);
1220 }
1221 self.decode = Some(self.fresh_decode_state());
1222 *self.gen_acc.borrow_mut() = GenerateProfile::default();
1223 let result = (|| {
1224 let t_pre = Instant::now();
1225 let mut logits = if self.can_batch_prefill() && tokens.len() > 1 {
1226 self.forward_prompt(&tokens)?
1227 } else {
1228 let mut last = Vec::new();
1229 for &tok in &tokens {
1230 last = self.forward_step(tok)?;
1231 }
1232 last
1233 };
1234 if self.profile_on {
1235 self.gen_acc.borrow_mut().prefill_ms = elapsed_ms(t_pre);
1236 }
1237 let mut generated = Vec::new();
1238 let t_dec = Instant::now();
1239 for _ in 0..opts.max_tokens {
1240 let next = argmax(&logits);
1242 generated.push(next);
1243 tokens.push(next);
1244 if self.is_stop_id(next) {
1245 generated.pop();
1246 break;
1247 }
1248 logits = self.forward_step(next)?;
1249 }
1250 if self.profile_on {
1251 self.gen_acc.borrow_mut().decode_ms = elapsed_ms(t_dec);
1252 }
1253 let text = self.decode_tokens(&generated);
1254 Ok(Generation {
1255 tokens: generated,
1256 text,
1257 })
1258 })();
1259 if self.profile_on {
1260 let mut p = self.last_profile.take().unwrap_or(EngineProfile {
1261 compute: self.compute_label.clone(),
1262 load: load_profile_take(),
1263 generate: None,
1264 ci_fail: false,
1265 });
1266 p.generate = Some(self.gen_acc.borrow().clone());
1267 self.last_profile = Some(p);
1268 }
1269 self.decode = None;
1270 result
1271 }
1272
1273 pub fn decode_tokens(&self, ids: &[u32]) -> String {
1276 match &self.tokenizer {
1277 Some(tok) => {
1278 let raw = tok.decode_opts(ids, false);
1279 strip_assistant_visible(&raw)
1280 }
1281 None => decode_placeholders(ids),
1282 }
1283 }
1284
1285 pub fn encode_text(&self, text: &str) -> Vec<u32> {
1287 match &self.tokenizer {
1288 Some(tok) => match tok.encode(text) {
1289 Ok(ids) if !ids.is_empty() => ids,
1290 Ok(_) => encode_naive(text, self.conf.vocab_size as u32),
1291 Err(_) => encode_naive(text, self.conf.vocab_size as u32),
1292 },
1293 None => encode_naive(text, self.conf.vocab_size as u32),
1294 }
1295 }
1296
1297 pub fn encode_chat(&self, messages: &[ChatTurn]) -> Vec<u32> {
1299 let family = if self.family.path().contains("gemma-4") {
1302 self.family.path()
1303 } else {
1304 self.tokenizer
1305 .as_ref()
1306 .and_then(|t| t.chat_family_hint())
1307 .unwrap_or(self.family.path())
1308 };
1309 let prompt = apply_chat_template(family, messages);
1310 self.encode_text(&prompt)
1311 }
1312
1313 fn is_stop_id(&self, id: u32) -> bool {
1314 match &self.tokenizer {
1315 Some(t) => t.is_stop(id),
1316 None => id == 0,
1317 }
1318 }
1319
1320 pub fn arch(&self) -> ArchClass {
1321 self.family.arch
1322 }
1323
1324 pub fn graph_hook_name(&self) -> &'static str {
1325 graph_hook(self.family.arch)
1326 }
1327
1328 pub fn embed_text(&self, text: &str) -> Result<Vec<f32>, EngineError> {
1330 let toks = self.encode_text(text);
1331 let hidden = self.conf.hidden_size;
1332 let vocab = self.conf.vocab_size;
1333 let mut acc = vec![0.0f32; hidden];
1334 if toks.is_empty() {
1335 return Ok(acc);
1336 }
1337 for &tok in &toks {
1338 let tid = (tok as usize) % vocab;
1339 let row = &self.weights.emb.data[tid * hidden..(tid + 1) * hidden];
1340 for i in 0..hidden {
1341 acc[i] += row[i];
1342 }
1343 }
1344 let inv = 1.0 / toks.len() as f32;
1345 for v in &mut acc {
1346 *v *= inv;
1347 }
1348 Ok(acc)
1349 }
1350
1351 pub fn vision_prefix(
1353 &self,
1354 rgb: &[u8],
1355 height: usize,
1356 width: usize,
1357 ) -> Result<Vec<f32>, EngineError> {
1358 if !matches!(self.family.arch, ArchClass::VL | ArchClass::VLA) {
1359 return Err(EngineError::Unsupported(format!(
1360 "vision_prefix not available for arch {:?}",
1361 self.family.arch
1362 )));
1363 }
1364 let Some(proj) = &self.weights.vision else {
1365 return Err(EngineError::Unsupported(format!(
1366 "{}: no vision projector tensor in bundle",
1367 self.family.path()
1368 )));
1369 };
1370 let hidden = self.conf.hidden_size;
1371 if hidden == 0 || proj.data.len() % hidden != 0 {
1372 return Err(EngineError::ShapeMismatch(
1373 "vision projector not divisible by hidden_size".into(),
1374 ));
1375 }
1376 let in_f = proj.data.len() / hidden;
1377 let need = height
1378 .checked_mul(width)
1379 .and_then(|n| n.checked_mul(3))
1380 .ok_or_else(|| EngineError::InvalidParam("vision size overflow".into()))?;
1381 if rgb.len() < need {
1382 return Err(EngineError::ShapeMismatch(format!(
1383 "rgb len {} < {}x{}x3",
1384 rgb.len(),
1385 height,
1386 width
1387 )));
1388 }
1389 let mut feat = vec![0.0f32; in_f];
1390 let pixels = height * width;
1391 if in_f == 3 {
1392 let mut acc = [0.0f32; 3];
1393 for p in 0..pixels {
1394 acc[0] += rgb[p * 3] as f32 / 255.0;
1395 acc[1] += rgb[p * 3 + 1] as f32 / 255.0;
1396 acc[2] += rgb[p * 3 + 2] as f32 / 255.0;
1397 }
1398 let s = 1.0 / pixels.max(1) as f32;
1399 feat[0] = acc[0] * s;
1400 feat[1] = acc[1] * s;
1401 feat[2] = acc[2] * s;
1402 } else {
1403 for i in 0..in_f {
1404 feat[i] = rgb[i % need] as f32 / 255.0;
1405 }
1406 }
1407 self.wmm(proj, &feat, hidden, in_f, GemmAcct::Other)
1408 }
1409
1410 pub fn predict_action(&self, prompt: &str, action_dim: usize) -> Result<Vec<f32>, EngineError> {
1412 if self.family.arch != ArchClass::VLA {
1413 return Err(EngineError::Unsupported(format!(
1414 "predict_action requires VLA, got {:?}",
1415 self.family.arch
1416 )));
1417 }
1418 if action_dim == 0 {
1419 return Err(EngineError::InvalidParam("action_dim must be > 0".into()));
1420 }
1421 let Some(head) = &self.weights.action else {
1422 return Err(EngineError::Unsupported(format!(
1423 "{}: no action head tensor in bundle",
1424 self.family.path()
1425 )));
1426 };
1427 let h = self.embed_text(prompt)?;
1428 let hidden = self.conf.hidden_size;
1429 if head.data.len() % hidden != 0 {
1430 return Err(EngineError::ShapeMismatch(
1431 "action head not divisible by hidden_size".into(),
1432 ));
1433 }
1434 let out_f = head.data.len() / hidden;
1435 if out_f != action_dim {
1436 return Err(EngineError::ShapeMismatch(format!(
1437 "action head out {out_f} != requested {action_dim}"
1438 )));
1439 }
1440 self.wmm(head, &h, out_f, hidden, GemmAcct::Other)
1441 }
1442
1443 pub fn transcribe_pcm16le(&self, pcm: &[u8]) -> Result<String, EngineError> {
1445 asr_transcribe_pcm16le(pcm, self.conf.vocab_size as u32)
1446 }
1447
1448 fn norm(&self, x: &[f32], weight: &[f32]) -> Result<Vec<f32>, EngineError> {
1449 if self.use_gemma_norm {
1450 rms_norm_gemma(x, weight, 1e-6)
1451 } else {
1452 rms_norm(x, weight, 1e-6)
1453 }
1454 }
1455
1456 fn add_normed_residual(
1457 &self,
1458 x: &mut [f32],
1459 y: &[f32],
1460 post_norm: Option<&[f32]>,
1461 ) -> Result<(), EngineError> {
1462 if y.len() != x.len() {
1463 return Err(EngineError::ShapeMismatch(
1464 "residual length mismatch".into(),
1465 ));
1466 }
1467 if let Some(w) = post_norm {
1468 let yn = self.norm(y, w)?;
1469 for (a, b) in x.iter_mut().zip(yn.iter()) {
1470 *a += *b;
1471 }
1472 } else {
1473 for (a, b) in x.iter_mut().zip(y.iter()) {
1474 *a += *b;
1475 }
1476 }
1477 Ok(())
1478 }
1479
1480 fn attn_scale(&self, head_dim: usize) -> f32 {
1481 if self.use_gemma4 {
1482 1.0
1483 } else {
1484 1.0 / (head_dim as f32).sqrt()
1485 }
1486 }
1487
1488 fn attn_window(&self, kind: AttnKind) -> Option<usize> {
1492 if kind != AttnKind::Sliding {
1493 return None;
1494 }
1495 self.conf.sliding_window.filter(|w| *w > 0)
1496 }
1497
1498 fn layer_rope_params(&self, kind: AttnKind) -> (f32, RopeMode) {
1499 if self.use_gemma4 {
1500 match kind {
1501 AttnKind::Sliding => (10_000.0, RopeMode::Full),
1502 AttnKind::Full => {
1503 let factor = self.conf.partial_rotary_factor.unwrap_or(1.0);
1505 let mode = if factor > 0.0 && factor < 1.0 {
1506 RopeMode::Proportional(factor)
1507 } else {
1508 RopeMode::Full
1509 };
1510 (1_000_000.0, mode)
1511 }
1512 }
1513 } else if self.family.path().contains("qwen3.5") {
1514 let factor = self.conf.partial_rotary_factor.unwrap_or(0.25);
1517 let mode = if factor > 0.0 && factor < 1.0 {
1518 RopeMode::Partial(factor)
1519 } else {
1520 RopeMode::Full
1521 };
1522 (self.conf.rope_theta, mode)
1523 } else {
1524 (self.conf.rope_theta, RopeMode::Full)
1525 }
1526 }
1527
1528 fn layer_head_dim(
1529 &self,
1530 kind: AttnKind,
1531 q_dim: usize,
1532 n_heads: usize,
1533 ) -> Result<usize, EngineError> {
1534 if n_heads == 0 || !q_dim.is_multiple_of(n_heads) {
1535 return Err(EngineError::ShapeMismatch(
1536 "q_dim not divisible by num_attention_heads".into(),
1537 ));
1538 }
1539 let configured = match kind {
1540 AttnKind::Full => self.conf.global_head_dim.or(self.conf.head_dim),
1541 AttnKind::Sliding => self.conf.head_dim,
1542 };
1543 Ok(configured
1544 .filter(|d| *d > 0 && q_dim == n_heads * *d)
1545 .unwrap_or(q_dim / n_heads))
1546 }
1547
1548 fn apply_rope(
1549 x: &mut [f32],
1550 head_dim: usize,
1551 pos: usize,
1552 theta: f32,
1553 mode: RopeMode,
1554 ) -> Result<(), EngineError> {
1555 match mode {
1556 RopeMode::Full => rope_half(x, head_dim, pos, theta),
1557 RopeMode::Proportional(factor) => {
1558 rope_half_proportional(x, head_dim, factor, pos, theta)
1561 }
1562 RopeMode::Partial(factor) => {
1563 let rotary_dim = (factor * head_dim as f32) as usize & !1;
1566 if rotary_dim < 2 || rotary_dim >= head_dim {
1567 rope_half(x, head_dim, pos, theta)
1568 } else {
1569 rope_half_partial(x, head_dim, rotary_dim, pos, theta)
1570 }
1571 }
1572 }
1573 }
1574
1575 fn apply_v_norm(
1576 &self,
1577 v: Vec<f32>,
1578 v_norm: Option<&[f32]>,
1579 head_dim: usize,
1580 ) -> Result<Vec<f32>, EngineError> {
1581 if let Some(vn) = v_norm {
1582 if vn.len() != head_dim {
1583 return Err(EngineError::ShapeMismatch(format!(
1584 "v_norm len {} != head_dim {head_dim}",
1585 vn.len()
1586 )));
1587 }
1588 self.norm(&v, vn)
1589 } else if self.use_gemma4 {
1590 let ones = vec![1.0f32; head_dim];
1591 rms_norm(&v, &ones, 1e-6)
1592 } else {
1593 Ok(v)
1594 }
1595 }
1596
1597 fn compute_ple_inputs(
1598 &self,
1599 toks: &[u32],
1600 embeds: &[f32],
1601 ) -> Result<Option<Vec<f32>>, EngineError> {
1602 let Some(ple) = &self.weights.ple else {
1603 return Ok(None);
1604 };
1605 let hidden = self.conf.hidden_size;
1606 let n_layers = self.weights.layers.len();
1607 let d = ple.d;
1608 let packed = n_layers * d;
1609 let seq = toks.len();
1610 if seq == 0 || embeds.len() != seq * hidden {
1611 return Err(EngineError::ShapeMismatch(
1612 "PLE embed sequence length mismatch".into(),
1613 ));
1614 }
1615 let scale_lookup = (d as f32).sqrt();
1616 let ple_vocab = ple.embed.len() / packed;
1617 if ple_vocab == 0 {
1618 return Err(EngineError::ShapeMismatch("PLE embed vocab is 0".into()));
1619 }
1620 let mut lookup = vec![0.0f32; seq * packed];
1621 for (t, &tok) in toks.iter().enumerate() {
1622 let tid = (tok as usize) % ple_vocab;
1623 let row = &ple.embed[tid * packed..(tid + 1) * packed];
1624 for i in 0..packed {
1625 lookup[t * packed + i] = row[i] * scale_lookup;
1626 }
1627 }
1628 let proj_scale = (hidden as f32).sqrt().recip();
1629 let mut proj = self.wmm(&ple.proj, embeds, packed, hidden, GemmAcct::Other)?;
1630 for v in &mut proj {
1631 *v *= proj_scale;
1632 }
1633 proj = rms_norm(&proj, &ple.proj_norm, 1e-6)?;
1634 let inv_sqrt2 = std::f32::consts::FRAC_1_SQRT_2;
1635 for i in 0..proj.len() {
1636 proj[i] = (proj[i] + lookup[i]) * inv_sqrt2;
1637 }
1638 Ok(Some(proj))
1639 }
1640
1641 fn apply_ple(
1642 &self,
1643 x: &mut [f32],
1644 layer: &LayerWeights,
1645 li: usize,
1646 ple_tok: Option<&[f32]>,
1647 hidden: usize,
1648 ) -> Result<(), EngineError> {
1649 let (Some(ple), Some(ple_tok)) = (&layer.ple, ple_tok) else {
1650 return Ok(());
1651 };
1652 let d = self
1653 .weights
1654 .ple
1655 .as_ref()
1656 .map(|p| p.d)
1657 .ok_or_else(|| EngineError::Format("layer PLE without model PLE".into()))?;
1658 let n_layers = self.weights.layers.len();
1659 let seq = x.len() / hidden;
1660 let gate_out = self.wmm(&ple.gate, x, d, hidden, GemmAcct::Ffn)?;
1661 let mut gated = vec![0.0f32; seq * d];
1662 for t in 0..seq {
1663 for i in 0..d {
1664 let g = gelu_pytorch_tanh(gate_out[t * d + i]);
1665 let p = ple_tok[t * n_layers * d + li * d + i];
1666 gated[t * d + i] = g * p;
1667 }
1668 }
1669 let proj = self.wmm(&ple.proj, &gated, hidden, d, GemmAcct::Ffn)?;
1670 let nrm = self.norm(&proj, &ple.post_norm)?;
1671 for (a, b) in x.iter_mut().zip(nrm.iter()) {
1672 *a += *b;
1673 }
1674 Ok(())
1675 }
1676
1677 fn apply_layer_scalar(x: &mut [f32], scale: f32) {
1678 if (scale - 1.0).abs() < 1e-8 {
1679 return;
1680 }
1681 for v in x {
1682 *v *= scale;
1683 }
1684 }
1685
1686 fn apply_ffn(
1687 &self,
1688 layer: &LayerWeights,
1689 xn2: &[f32],
1690 hidden: usize,
1691 ) -> Result<Vec<f32>, EngineError> {
1692 match &layer.ffn {
1693 FfnWeights::Dense { gate, up, down } => {
1694 if gate.data.len() % hidden != 0 {
1695 return Err(EngineError::ShapeMismatch(
1696 "dense gate len not divisible by hidden".into(),
1697 ));
1698 }
1699 let inter = gate.data.len() / hidden;
1700 if inter == 0
1701 || up.data.len() != inter * hidden
1702 || down.data.len() != hidden * inter
1703 {
1704 return Err(EngineError::ShapeMismatch(
1705 "dense FFN weight shape mismatch".into(),
1706 ));
1707 }
1708 let g = self.wmm(gate, xn2, inter, hidden, GemmAcct::Ffn)?;
1709 let u = self.wmm(up, xn2, inter, hidden, GemmAcct::Ffn)?;
1710 let h = if self.use_geglu {
1711 geglu(&g, &u)?
1712 } else {
1713 swiglu(&g, &u)?
1714 };
1715 self.wmm(down, &h, hidden, inter, GemmAcct::Ffn)
1716 }
1717 FfnWeights::MoE {
1718 router,
1719 experts,
1720 top_k,
1721 use_sigmoid,
1722 } => {
1723 let n_exp = experts.len();
1724 let logits = self.wmm(router, xn2, n_exp, hidden, GemmAcct::Ffn)?;
1725 let (ids, weights) = moe_topk_route(&logits, *top_k, *use_sigmoid)?;
1726 let mut acc = vec![0.0f32; hidden];
1727 for (ei, &w) in ids.iter().zip(weights.iter()) {
1728 let ex = &experts[*ei];
1729 if ex.gate.data.len() % hidden != 0 {
1730 return Err(EngineError::ShapeMismatch(
1731 "expert gate len not divisible by hidden".into(),
1732 ));
1733 }
1734 let inter = ex.gate.data.len() / hidden;
1735 let g = self.wmm(&ex.gate, xn2, inter, hidden, GemmAcct::Ffn)?;
1736 let u = self.wmm(&ex.up, xn2, inter, hidden, GemmAcct::Ffn)?;
1737 let h = swiglu(&g, &u)?;
1738 let down = self.wmm(&ex.down, &h, hidden, inter, GemmAcct::Ffn)?;
1739 for i in 0..hidden {
1740 acc[i] += w * down[i];
1741 }
1742 }
1743 Ok(acc)
1744 }
1745 }
1746 }
1747
1748 fn fresh_decode_state(&self) -> DecodeState {
1749 let hidden = self.conf.hidden_size;
1750 DecodeState {
1751 k_caches: (0..self.conf.num_layers).map(|_| Vec::new()).collect(),
1752 v_caches: (0..self.conf.num_layers).map(|_| Vec::new()).collect(),
1753 last_kv_src: HashMap::new(),
1754 conv_states: self
1755 .weights
1756 .layers
1757 .iter()
1758 .map(|layer| match &layer.op {
1759 LayerOp::Conv(c) => {
1760 let hist = c.kernel_size.saturating_sub(1);
1761 Some(vec![0.0f32; hidden * hist])
1762 }
1763 LayerOp::Linear(d) => {
1764 let conv_dim = d.n_k_heads * d.head_k * 2 + d.n_v_heads * d.head_v;
1765 let hist = d.conv_k.saturating_sub(1);
1766 Some(vec![0.0f32; conv_dim * hist])
1767 }
1768 LayerOp::Attn(_) => None,
1769 })
1770 .collect(),
1771 delta_states: self
1772 .weights
1773 .layers
1774 .iter()
1775 .map(|layer| match &layer.op {
1776 LayerOp::Linear(d) => Some(vec![0.0f32; d.n_v_heads * d.head_k * d.head_v]),
1777 _ => None,
1778 })
1779 .collect(),
1780 pos: 0,
1781 }
1782 }
1783
1784 #[cfg(test)]
1786 fn forward(&self, tokens: &[u32]) -> Result<Vec<f32>, EngineError> {
1787 let mut state = self.fresh_decode_state();
1788 let mut logits = Vec::new();
1789 for &tok in tokens {
1790 logits = self.forward_step_with(&mut state, tok)?;
1791 }
1792 Ok(logits)
1793 }
1794
1795 fn forward_prompt(&mut self, toks: &[u32]) -> Result<Vec<f32>, EngineError> {
1796 let mut owned = self.decode.take().ok_or_else(|| {
1797 EngineError::InvalidParam("decode state missing; call generate/prefill first".into())
1798 })?;
1799 let logits = self.forward_prompt_with(&mut owned, toks);
1800 self.decode = Some(owned);
1801 logits
1802 }
1803
1804 fn apply_rope_seq(
1805 x: &mut [f32],
1806 seq: usize,
1807 tok_dim: usize,
1808 head_dim: usize,
1809 pos0: usize,
1810 theta: f32,
1811 mode: RopeMode,
1812 ) -> Result<(), EngineError> {
1813 if x.len() != seq * tok_dim {
1814 return Err(EngineError::ShapeMismatch(
1815 "rope seq buffer length mismatch".into(),
1816 ));
1817 }
1818 for t in 0..seq {
1819 Self::apply_rope(
1820 &mut x[t * tok_dim..(t + 1) * tok_dim],
1821 head_dim,
1822 pos0 + t,
1823 theta,
1824 mode,
1825 )?;
1826 }
1827 Ok(())
1828 }
1829
1830 fn forward_prompt_with(
1831 &self,
1832 state: &mut DecodeState,
1833 toks: &[u32],
1834 ) -> Result<Vec<f32>, EngineError> {
1835 if toks.is_empty() {
1836 return Err(EngineError::InvalidParam("empty prompt".into()));
1837 }
1838 let hidden = self.conf.hidden_size;
1839 let n_heads = self.conf.num_attention_heads;
1840 let n_kv = self.conf.num_kv_heads;
1841 let vocab = self.conf.vocab_size;
1842 let seq = toks.len();
1843 if hidden == 0 {
1844 return Err(EngineError::ShapeMismatch("hidden_size is 0".into()));
1845 }
1846 if self.weights.emb.data.len() < vocab.saturating_mul(hidden)
1847 || !self.weights.emb.data.len().is_multiple_of(hidden)
1848 {
1849 return Err(EngineError::ShapeMismatch(format!(
1850 "embedding length {} not compatible with vocab={vocab} hidden={hidden}",
1851 self.weights.emb.data.len()
1852 )));
1853 }
1854 let pos0 = state.pos;
1855 let mut x = vec![0.0f32; seq * hidden];
1856 for (t, &tok) in toks.iter().enumerate() {
1857 let tid = (tok as usize) % vocab;
1858 x[t * hidden..(t + 1) * hidden]
1859 .copy_from_slice(&self.weights.emb.data[tid * hidden..(tid + 1) * hidden]);
1860 }
1861 if self.embed_scale != 1.0 {
1862 for v in &mut x {
1863 *v *= self.embed_scale;
1864 }
1865 }
1866 let ple_tok = self.compute_ple_inputs(toks, &x)?;
1867
1868 for (li, layer) in self.weights.layers.iter().enumerate() {
1869 let xn = self.norm(&x, &layer.attn_norm)?;
1870 match &layer.op {
1871 LayerOp::Attn(attn) => {
1872 if attn.wq.data.len() % hidden != 0 {
1873 return Err(EngineError::ShapeMismatch(
1874 "attn q proj weight not divisible by hidden_size".into(),
1875 ));
1876 }
1877 let q_out = attn.wq.data.len() / hidden;
1878 let q_dim = if attn.q_gate { q_out / 2 } else { q_out };
1879 let head_dim = self.layer_head_dim(attn.kind, q_dim, n_heads)?;
1880 if attn.wo.data.len() != hidden * q_dim {
1881 return Err(EngineError::ShapeMismatch(format!(
1882 "attn output proj weight shape mismatch (wo_len={} hidden={hidden} q_dim={q_dim})",
1883 attn.wo.data.len()
1884 )));
1885 }
1886 let mixed = self.wmm(&attn.wq, &xn, q_out, hidden, GemmAcct::Attn)?;
1887 let (mut q, gate) = if attn.q_gate {
1888 split_interleaved_q_gate(&mixed, seq, n_heads, head_dim)?
1889 } else {
1890 (mixed, Vec::new())
1891 };
1892 if let Some(qn) = &attn.q_norm {
1893 if qn.len() != head_dim {
1894 return Err(EngineError::ShapeMismatch(format!(
1895 "q_norm len {} != head_dim {head_dim}",
1896 qn.len()
1897 )));
1898 }
1899 q = self.norm(&q, qn)?;
1900 }
1901 let (theta, rope) = self.layer_rope_params(attn.kind);
1902 Self::apply_rope_seq(&mut q, seq, q_dim, head_dim, pos0, theta, rope)?;
1903
1904 let (k_src, v_src) = if let (Some(wk), Some(wv)) = (&attn.wk, &attn.wv) {
1905 if wk.data.len() % hidden != 0 || wv.data.len() % hidden != 0 {
1906 return Err(EngineError::ShapeMismatch(
1907 "attn kv proj weight not divisible by hidden_size".into(),
1908 ));
1909 }
1910 let k_dim = wk.data.len() / hidden;
1911 let v_dim = wv.data.len() / hidden;
1912 if k_dim != n_kv * head_dim || v_dim != n_kv * head_dim {
1913 return Err(EngineError::ShapeMismatch(format!(
1914 "kv dims {k_dim}/{v_dim} != n_kv*head_dim {}",
1915 n_kv * head_dim
1916 )));
1917 }
1918 let mut k = self.wmm(wk, &xn, k_dim, hidden, GemmAcct::Attn)?;
1919 let mut v = self.wmm(wv, &xn, v_dim, hidden, GemmAcct::Attn)?;
1920 if let Some(kn) = &attn.k_norm {
1921 if kn.len() != head_dim {
1922 return Err(EngineError::ShapeMismatch(format!(
1923 "k_norm len {} != head_dim {head_dim}",
1924 kn.len()
1925 )));
1926 }
1927 k = self.norm(&k, kn)?;
1928 }
1929 Self::apply_rope_seq(
1930 &mut k,
1931 seq,
1932 k_dim,
1933 head_dim,
1934 pos0,
1935 theta,
1936 rope,
1937 )?;
1938 v = self.apply_v_norm(v, attn.v_norm.as_deref(), head_dim)?;
1939 state.k_caches[li] = k;
1940 state.v_caches[li] = v;
1941 state.last_kv_src.insert(attn.kind, li);
1942 (li, li)
1943 } else {
1944 let src = state.last_kv_src.get(&attn.kind).copied().ok_or_else(|| {
1945 EngineError::Format(format!(
1946 "KV-consumer layer {li} has no producer of kind {:?}",
1947 attn.kind
1948 ))
1949 })?;
1950 (src, src)
1951 };
1952 let attn_out = attention_causal_with_scale(
1953 &q,
1954 &state.k_caches[k_src],
1955 &state.v_caches[v_src],
1956 n_heads,
1957 n_kv,
1958 head_dim,
1959 self.attn_scale(head_dim),
1960 self.attn_window(attn.kind),
1961 )?;
1962 let mut attn_out = attn_out;
1963 if attn.q_gate {
1964 apply_sigmoid_gate(&mut attn_out, &gate)?;
1965 }
1966 let ao = self.wmm(&attn.wo, &attn_out, hidden, q_dim, GemmAcct::Attn)?;
1967 self.add_normed_residual(&mut x, &ao, layer.post_attn_norm.as_deref())?;
1968 }
1969 LayerOp::Conv(_) | LayerOp::Linear(_) => {
1970 return Err(EngineError::Unsupported(
1971 "batched prefill is only implemented for attention+dense FFN layers"
1972 .into(),
1973 ));
1974 }
1975 }
1976 let xn2 = self.norm(&x, &layer.ffn_norm)?;
1977 let down = self.apply_ffn(layer, &xn2, hidden)?;
1978 self.add_normed_residual(&mut x, &down, layer.post_ffn_norm.as_deref())?;
1979 self.apply_ple(&mut x, layer, li, ple_tok.as_deref(), hidden)?;
1980 Self::apply_layer_scalar(&mut x, layer.layer_scalar);
1981 }
1982 state.pos = pos0 + seq;
1983 let last = &x[(seq - 1) * hidden..seq * hidden];
1984 let xn = self.norm(last, &self.weights.output_norm)?;
1985 if !self.weights.output.data.len().is_multiple_of(hidden) {
1986 return Err(EngineError::ShapeMismatch(format!(
1987 "lm_head len {} not divisible by hidden {hidden}",
1988 self.weights.output.data.len()
1989 )));
1990 }
1991 let out_rows = self.weights.output.data.len() / hidden;
1992 let logits = self.wmm(
1993 &self.weights.output,
1994 &xn,
1995 out_rows,
1996 hidden,
1997 GemmAcct::LmHead,
1998 )?;
1999 Ok(self.softcap_logits(logits))
2000 }
2001
2002 fn forward_step(&mut self, tok: u32) -> Result<Vec<f32>, EngineError> {
2003 let mut owned = self.decode.take().ok_or_else(|| {
2004 EngineError::InvalidParam("decode state missing; call generate/prefill first".into())
2005 })?;
2006 let logits = self.forward_step_with(&mut owned, tok);
2007 self.decode = Some(owned);
2008 logits
2009 }
2010
2011 fn forward_step_with(
2012 &self,
2013 state: &mut DecodeState,
2014 tok: u32,
2015 ) -> Result<Vec<f32>, EngineError> {
2016 let hidden = self.conf.hidden_size;
2017 let n_heads = self.conf.num_attention_heads;
2018 let n_kv = self.conf.num_kv_heads;
2019 let vocab = self.conf.vocab_size;
2020 if hidden == 0 {
2021 return Err(EngineError::ShapeMismatch("hidden_size is 0".into()));
2022 }
2023 if self.weights.emb.data.len() < vocab.saturating_mul(hidden)
2024 || !self.weights.emb.data.len().is_multiple_of(hidden)
2025 {
2026 return Err(EngineError::ShapeMismatch(format!(
2027 "embedding length {} not compatible with vocab={vocab} hidden={hidden}",
2028 self.weights.emb.data.len()
2029 )));
2030 }
2031 let pos = state.pos;
2032 let tid = (tok as usize) % vocab;
2033 let mut x = vec![0.0f32; hidden];
2034 x.copy_from_slice(&self.weights.emb.data[tid * hidden..(tid + 1) * hidden]);
2035 if self.embed_scale != 1.0 {
2036 for v in &mut x {
2037 *v *= self.embed_scale;
2038 }
2039 }
2040 let ple_tok = self.compute_ple_inputs(&[tok], &x)?;
2041
2042 for (li, layer) in self.weights.layers.iter().enumerate() {
2043 let xn = self.norm(&x, &layer.attn_norm)?;
2044 match &layer.op {
2045 LayerOp::Attn(attn) => {
2046 if attn.wq.data.len() % hidden != 0 {
2047 return Err(EngineError::ShapeMismatch(
2048 "attn q proj weight not divisible by hidden_size".into(),
2049 ));
2050 }
2051 let q_out = attn.wq.data.len() / hidden;
2052 let q_dim = if attn.q_gate { q_out / 2 } else { q_out };
2053 let head_dim = self.layer_head_dim(attn.kind, q_dim, n_heads)?;
2054 if attn.wo.data.len() != hidden * q_dim {
2055 return Err(EngineError::ShapeMismatch(format!(
2056 "attn output proj weight shape mismatch (wo_len={} hidden={hidden} q_dim={q_dim})",
2057 attn.wo.data.len()
2058 )));
2059 }
2060 let mixed = self.wmm(&attn.wq, &xn, q_out, hidden, GemmAcct::Attn)?;
2061 let (mut q, gate) = if attn.q_gate {
2062 split_interleaved_q_gate(&mixed, 1, n_heads, head_dim)?
2063 } else {
2064 (mixed, Vec::new())
2065 };
2066 if let Some(qn) = &attn.q_norm {
2067 if qn.len() != head_dim {
2068 return Err(EngineError::ShapeMismatch(format!(
2069 "q_norm len {} != head_dim {head_dim}",
2070 qn.len()
2071 )));
2072 }
2073 q = self.norm(&q, qn)?;
2074 }
2075 let (theta, rope) = self.layer_rope_params(attn.kind);
2076 Self::apply_rope(&mut q, head_dim, pos, theta, rope)?;
2077
2078 let (k_src, v_src) = if let (Some(wk), Some(wv)) = (&attn.wk, &attn.wv) {
2079 if wk.data.len() % hidden != 0 || wv.data.len() % hidden != 0 {
2080 return Err(EngineError::ShapeMismatch(
2081 "attn kv proj weight not divisible by hidden_size".into(),
2082 ));
2083 }
2084 let k_dim = wk.data.len() / hidden;
2085 let v_dim = wv.data.len() / hidden;
2086 if k_dim != n_kv * head_dim || v_dim != n_kv * head_dim {
2087 return Err(EngineError::ShapeMismatch(format!(
2088 "kv dims {k_dim}/{v_dim} != n_kv*head_dim {}",
2089 n_kv * head_dim
2090 )));
2091 }
2092 let mut k = self.wmm(wk, &xn, k_dim, hidden, GemmAcct::Attn)?;
2093 let mut v = self.wmm(wv, &xn, v_dim, hidden, GemmAcct::Attn)?;
2094 if let Some(kn) = &attn.k_norm {
2095 if kn.len() != head_dim {
2096 return Err(EngineError::ShapeMismatch(format!(
2097 "k_norm len {} != head_dim {head_dim}",
2098 kn.len()
2099 )));
2100 }
2101 k = self.norm(&k, kn)?;
2102 }
2103 Self::apply_rope(&mut k, head_dim, pos, theta, rope)?;
2104 v = self.apply_v_norm(v, attn.v_norm.as_deref(), head_dim)?;
2105 state.k_caches[li].extend_from_slice(&k);
2106 state.v_caches[li].extend_from_slice(&v);
2107 state.last_kv_src.insert(attn.kind, li);
2108 (li, li)
2109 } else {
2110 let src = state.last_kv_src.get(&attn.kind).copied().ok_or_else(|| {
2111 EngineError::Format(format!(
2112 "KV-consumer layer {li} has no producer of kind {:?}",
2113 attn.kind
2114 ))
2115 })?;
2116 (src, src)
2117 };
2118 let kv_dim = n_kv * head_dim;
2119 let (k_view, v_view) = kv_sliding_view(
2120 &state.k_caches[k_src],
2121 &state.v_caches[v_src],
2122 kv_dim,
2123 self.attn_window(attn.kind),
2124 )?;
2125 let attn_out = attention_with_scale(
2126 &q,
2127 k_view,
2128 v_view,
2129 n_heads,
2130 n_kv,
2131 head_dim,
2132 self.attn_scale(head_dim),
2133 )?;
2134 let mut attn_out = attn_out;
2135 if attn.q_gate {
2136 apply_sigmoid_gate(&mut attn_out, &gate)?;
2137 }
2138 let ao = self.wmm(&attn.wo, &attn_out, hidden, q_dim, GemmAcct::Attn)?;
2139 self.add_normed_residual(&mut x, &ao, layer.post_attn_norm.as_deref())?;
2140 }
2141 LayerOp::Conv(conv) => {
2142 let bcx = self.wmm(&conv.in_proj, &xn, 3 * hidden, hidden, GemmAcct::Attn)?;
2143 let mut bx = vec![0.0f32; hidden];
2144 let mut c_gate = vec![0.0f32; hidden];
2145 for i in 0..hidden {
2146 let b = bcx[i];
2147 let c = bcx[hidden + i];
2148 let xx = bcx[2 * hidden + i];
2149 bx[i] = b * xx;
2150 c_gate[i] = c;
2151 }
2152 let cstate = state.conv_states[li]
2153 .as_mut()
2154 .ok_or_else(|| EngineError::ShapeMismatch("missing conv state".into()))?;
2155 let conv_y =
2156 short_conv_step(&bx, &conv.kernel, cstate, hidden, conv.kernel_size)?;
2157 let mut y = vec![0.0f32; hidden];
2158 for i in 0..hidden {
2159 y[i] = c_gate[i] * conv_y[i];
2160 }
2161 let ao = self.wmm(&conv.out_proj, &y, hidden, hidden, GemmAcct::Attn)?;
2162 self.add_normed_residual(&mut x, &ao, layer.post_attn_norm.as_deref())?;
2163 }
2164 LayerOp::Linear(dn) => {
2165 let key_dim = dn.n_k_heads * dn.head_k;
2166 let value_dim = dn.n_v_heads * dn.head_v;
2167 let qkvz_out = 2 * key_dim + 2 * value_dim;
2168 let mixed = self.wmm(&dn.qkvz, &xn, qkvz_out, hidden, GemmAcct::Attn)?;
2169 let mut q = mixed[0..key_dim].to_vec();
2170 let mut k = mixed[key_dim..2 * key_dim].to_vec();
2171 let mut v = mixed[2 * key_dim..2 * key_dim + value_dim].to_vec();
2172 let z = mixed[2 * key_dim + value_dim..].to_vec();
2173 let mut qkv = Vec::with_capacity(key_dim * 2 + value_dim);
2174 qkv.extend_from_slice(&q);
2175 qkv.extend_from_slice(&k);
2176 qkv.extend_from_slice(&v);
2177 let conv_dim = qkv.len();
2178 let cstate = state.conv_states[li].as_mut().ok_or_else(|| {
2179 EngineError::ShapeMismatch("missing delta conv state".into())
2180 })?;
2181 let mut mixed_c = short_conv_step(&qkv, &dn.conv, cstate, conv_dim, dn.conv_k)?;
2182 silu_vec(&mut mixed_c);
2183 q.copy_from_slice(&mixed_c[0..key_dim]);
2184 k.copy_from_slice(&mixed_c[key_dim..2 * key_dim]);
2185 v.copy_from_slice(&mixed_c[2 * key_dim..]);
2186 let ba = self.wmm(&dn.ba, &xn, 2 * dn.n_v_heads, hidden, GemmAcct::Attn)?;
2187 let mut beta = vec![0.0f32; dn.n_v_heads];
2188 let mut g = vec![0.0f32; dn.n_v_heads];
2189 for h in 0..dn.n_v_heads {
2190 beta[h] = 1.0 / (1.0 + (-ba[h]).exp());
2191 let alpha =
2192 -dn.a_log[h].exp() * softplus(ba[dn.n_v_heads + h] + dn.dt_bias[h]);
2193 g[h] = alpha.exp();
2194 }
2195 if dn.n_v_heads != dn.n_k_heads {
2196 return Err(EngineError::Unsupported(
2197 "DeltaNet GQA (n_v != n_k) not implemented".into(),
2198 ));
2199 }
2200 let s = state.delta_states[li].as_mut().ok_or_else(|| {
2201 EngineError::ShapeMismatch("missing delta recurrent state".into())
2202 })?;
2203 let mut core = gated_delta_step(GatedDeltaStep {
2204 q: &q,
2205 k: &k,
2206 v: &v,
2207 g: &g,
2208 beta: &beta,
2209 state: s,
2210 n_heads: dn.n_v_heads,
2211 dk: dn.head_k,
2212 dv: dn.head_v,
2213 })?;
2214 core = rms_norm(&core, &dn.out_norm, 1e-6)?;
2216 let mut z_act = z;
2217 silu_vec(&mut z_act);
2218 for i in 0..core.len() {
2219 core[i] *= z_act[i];
2220 }
2221 let ao = self.wmm(&dn.out_proj, &core, hidden, value_dim, GemmAcct::Attn)?;
2222 self.add_normed_residual(&mut x, &ao, layer.post_attn_norm.as_deref())?;
2223 }
2224 }
2225 let xn2 = self.norm(&x, &layer.ffn_norm)?;
2226 let down = self.apply_ffn(layer, &xn2, hidden)?;
2227 self.add_normed_residual(&mut x, &down, layer.post_ffn_norm.as_deref())?;
2228 self.apply_ple(&mut x, layer, li, ple_tok.as_deref(), hidden)?;
2229 Self::apply_layer_scalar(&mut x, layer.layer_scalar);
2230 }
2231 state.pos += 1;
2232 let xn = self.norm(&x, &self.weights.output_norm)?;
2233 if !self.weights.output.data.len().is_multiple_of(hidden) {
2234 return Err(EngineError::ShapeMismatch(format!(
2235 "lm_head len {} not divisible by hidden {hidden}",
2236 self.weights.output.data.len()
2237 )));
2238 }
2239 let out_rows = self.weights.output.data.len() / hidden;
2240 let logits = self.wmm(&self.weights.output, &xn, out_rows, hidden, GemmAcct::LmHead)?;
2241 Ok(self.softcap_logits(logits))
2242 }
2243
2244 fn softcap_logits(&self, mut logits: Vec<f32>) -> Vec<f32> {
2245 if let Some(cap) = self.final_logit_softcap.filter(|c| *c > 0.0) {
2246 for x in &mut logits {
2247 *x = (*x / cap).tanh() * cap;
2248 }
2249 }
2250 logits
2251 }
2252}
2253
2254fn argmax(v: &[f32]) -> u32 {
2255 let mut best = 0usize;
2256 let mut best_v = f32::NEG_INFINITY;
2257 for (i, &x) in v.iter().enumerate() {
2258 if x > best_v {
2259 best_v = x;
2260 best = i;
2261 }
2262 }
2263 best as u32
2264}
2265
2266pub fn confidence_from_logits(logits: &[f32]) -> f32 {
2268 if logits.is_empty() {
2269 return 0.0;
2270 }
2271 let m = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
2272 let mut sum = 0.0f32;
2273 let mut maxp = 0.0f32;
2274 for &x in logits {
2275 let e = (x - m).exp();
2276 sum += e;
2277 if e > maxp {
2278 maxp = e;
2279 }
2280 }
2281 if sum > 0.0 {
2282 maxp / sum
2283 } else {
2284 0.0
2285 }
2286}
2287
2288#[allow(dead_code)]
2289pub fn cache_shapes_ok(cache: &HashMap<usize, Vec<f32>>, kv_dim: usize) -> bool {
2290 cache.values().all(|v| v.len().is_multiple_of(kv_dim))
2291}
2292
2293#[cfg(test)]
2294mod tests {
2295 use super::*;
2296 use crate::family::{arch_class_representatives, graph_hook, lookup_family, require_stage_b};
2297 use crate::fixture::write_tiny_q4_bundle;
2298 use aria_kernel::{resolve_compute, ComputePref};
2299 use serde_json::{json, Value};
2300
2301 #[test]
2302 fn gemma4_fills_hub_bundle_missing_geometry_fields() {
2303 let dir = tempfile::tempdir().unwrap();
2304 write_tiny_q4_bundle(dir.path()).unwrap();
2305 let cfg_path = dir.path().join("config.json");
2306 let raw = std::fs::read_to_string(&cfg_path).unwrap();
2307 let mut cfg: Value = serde_json::from_str(&raw).unwrap();
2308 let model = cfg["model"].as_object_mut().unwrap();
2310 for key in [
2311 "layer_types",
2312 "sliding_window",
2313 "partial_rotary_factor",
2314 "global_head_dim",
2315 "head_dim",
2316 ] {
2317 model.remove(key);
2318 }
2319 std::fs::write(&cfg_path, cfg.to_string()).unwrap();
2320 let s = SessionBuilder::new()
2321 .model(dir.path())
2322 .family("gemma/gemma-4-e2b-it")
2323 .build()
2324 .unwrap();
2325 assert_eq!(s.config().sliding_window, Some(512));
2326 assert_eq!(s.config().partial_rotary_factor, Some(0.25));
2327 assert!(s.config().head_dim.unwrap_or(0) > 0);
2328 assert!(s.config().global_head_dim.unwrap_or(0) > 0);
2329 assert_eq!(
2330 s.config().layer_types.as_ref().map(|t| t.len()),
2331 Some(s.config().num_layers)
2332 );
2333 }
2334
2335 #[test]
2336 fn generate_tokens() {
2337 let dir = tempfile::tempdir().unwrap();
2338 write_tiny_q4_bundle(dir.path()).unwrap();
2339 let mut s = SessionBuilder::new()
2340 .model(dir.path())
2341 .family("gemma/gemma-4-e2b-it")
2342 .build()
2343 .unwrap();
2344 assert_eq!(s.config().sliding_window, Some(512));
2345 assert_eq!(s.config().partial_rotary_factor, Some(0.25));
2346 assert_eq!(s.config().head_dim, Some(16));
2347 assert_eq!(s.config().global_head_dim, Some(16));
2348 assert_eq!(
2349 s.config().layer_types,
2350 Some(vec!["full_attention".into(), "full_attention".into()])
2351 );
2352 assert_eq!(
2353 s.layer_rope_params(AttnKind::Full),
2354 (1_000_000.0, RopeMode::Proportional(0.25))
2355 );
2356 assert_eq!(
2357 s.layer_rope_params(AttnKind::Sliding),
2358 (10_000.0, RopeMode::Full)
2359 );
2360 assert_eq!(s.attn_window(AttnKind::Sliding), Some(512));
2361 assert_eq!(s.attn_window(AttnKind::Full), None);
2362 let prompt = s.encode_text("hi");
2363 let gen = s
2364 .generate(
2365 &prompt,
2366 &GenerateOpts {
2367 max_tokens: 4,
2368 temperature: 0.0,
2369 },
2370 )
2371 .unwrap();
2372 assert!(!gen.tokens.is_empty());
2373 assert!(!gen.text.is_empty());
2374 }
2375
2376 #[test]
2377 fn split_interleaved_q_gate_matches_hf_chunk() {
2378 let mixed = vec![1.0, 2.0, 10.0, 20.0, 3.0, 4.0, 30.0, 40.0];
2380 let (q, g) = split_interleaved_q_gate(&mixed, 1, 2, 2).unwrap();
2381 assert_eq!(q, vec![1.0, 2.0, 3.0, 4.0]);
2382 assert_eq!(g, vec![10.0, 20.0, 30.0, 40.0]);
2383 }
2384
2385 #[test]
2386 fn qwen35_attn_output_gate_generate() {
2387 let dir = tempfile::tempdir().unwrap();
2388 let hidden = 8usize;
2389 let inter = 16usize;
2390 let vocab = 16usize;
2391 let n_heads = 2usize;
2392 let n_kv = 1usize;
2393 let head_dim = 4usize;
2394 let q_dim = n_heads * head_dim;
2395 let k_dim = n_kv * head_dim;
2396
2397 let mut tensors = serde_json::Map::new();
2398 let mut bin = Vec::new();
2399 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
2400 let offset = bin.len();
2401 for &v in data {
2402 bin.extend_from_slice(&v.to_le_bytes());
2403 }
2404 let nbytes = data.len() * 4;
2405 let mut meta = serde_json::Map::new();
2406 meta.insert("kind".into(), json!("raw"));
2407 meta.insert("dtype".into(), json!("f32"));
2408 meta.insert("shape".into(), json!(shape));
2409 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
2410 tensors.insert(name.to_string(), Value::Object(meta));
2411 };
2412 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
2413 add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
2414 let n1 = vec![1.0f32; hidden];
2415 add_raw("model.layers.0.input_layernorm.weight", vec![hidden], &n1);
2416 add_raw(
2417 "model.layers.0.post_attention_layernorm.weight",
2418 vec![hidden],
2419 &n1,
2420 );
2421 let wq = vec![0.02f32; 2 * q_dim * hidden];
2422 let wk = vec![0.02f32; k_dim * hidden];
2423 let wv = vec![0.02f32; k_dim * hidden];
2424 let wo = vec![0.02f32; hidden * q_dim];
2425 add_raw(
2426 "model.layers.0.self_attn.q_proj.weight",
2427 vec![2 * q_dim, hidden],
2428 &wq,
2429 );
2430 add_raw(
2431 "model.layers.0.self_attn.k_proj.weight",
2432 vec![k_dim, hidden],
2433 &wk,
2434 );
2435 add_raw(
2436 "model.layers.0.self_attn.v_proj.weight",
2437 vec![k_dim, hidden],
2438 &wv,
2439 );
2440 add_raw(
2441 "model.layers.0.self_attn.o_proj.weight",
2442 vec![hidden, q_dim],
2443 &wo,
2444 );
2445 let qn = vec![1.0f32; head_dim];
2446 add_raw("model.layers.0.self_attn.q_norm.weight", vec![head_dim], &qn);
2447 add_raw("model.layers.0.self_attn.k_norm.weight", vec![head_dim], &qn);
2448 let g = vec![0.02f32; inter * hidden];
2449 let d = vec![0.02f32; hidden * inter];
2450 add_raw(
2451 "model.layers.0.mlp.gate_proj.weight",
2452 vec![inter, hidden],
2453 &g,
2454 );
2455 add_raw("model.layers.0.mlp.up_proj.weight", vec![inter, hidden], &g);
2456 add_raw(
2457 "model.layers.0.mlp.down_proj.weight",
2458 vec![hidden, inter],
2459 &d,
2460 );
2461 add_raw("model.norm.weight", vec![hidden], &n1);
2462 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
2463 let cfg = json!({
2464 "format": "aria-quant-bundle",
2465 "format_version": 2,
2466 "quantization": "test",
2467 "hadamard_seed": 0,
2468 "model": {
2469 "hidden_size": hidden,
2470 "num_layers": 1,
2471 "num_attention_heads": n_heads,
2472 "num_kv_heads": n_kv,
2473 "head_dim": head_dim,
2474 "intermediate_size": inter,
2475 "vocab_size": vocab,
2476 "context_length": 32,
2477 "rope_theta": 10000.0,
2478 "layer_types": ["full_attention"]
2479 },
2480 "tensors": tensors
2481 });
2482 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
2483 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
2484 let mut s = SessionBuilder::new()
2485 .model(dir.path())
2486 .family("qwen/qwen3-0.6b")
2487 .build()
2488 .unwrap();
2489 let gen = s
2490 .generate(
2491 &[1, 2],
2492 &GenerateOpts {
2493 max_tokens: 2,
2494 temperature: 0.0,
2495 },
2496 )
2497 .unwrap();
2498 assert_eq!(gen.tokens.len(), 2);
2499 }
2500
2501 #[test]
2502 fn materialize_accepts_hf_tensor_names() {
2503 let dir = tempfile::tempdir().unwrap();
2505 let hidden = 8usize;
2506 let layers = 1usize;
2507 let inter = 16usize;
2508 let vocab = 16usize;
2509 let n_heads = 2usize;
2510 let n_kv = 1usize;
2511 let head_dim = 4usize; let q_dim = n_heads * head_dim;
2513 let k_dim = n_kv * head_dim;
2514
2515 let mut tensors = serde_json::Map::new();
2516 let mut bin = Vec::new();
2517 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
2518 let offset = bin.len();
2519 for &v in data {
2520 bin.extend_from_slice(&v.to_le_bytes());
2521 }
2522 let nbytes = data.len() * 4;
2523 let mut meta = serde_json::Map::new();
2524 meta.insert("kind".into(), json!("raw"));
2525 meta.insert("dtype".into(), json!("f32"));
2526 meta.insert("shape".into(), json!(shape));
2527 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
2528 tensors.insert(name.to_string(), Value::Object(meta));
2529 };
2530 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
2531 add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
2532 let n1 = vec![1.0f32; hidden];
2533 add_raw("model.layers.0.input_layernorm.weight", vec![hidden], &n1);
2534 add_raw(
2535 "model.layers.0.post_attention_layernorm.weight",
2536 vec![hidden],
2537 &n1,
2538 );
2539 let wq = vec![0.01f32; q_dim * hidden];
2540 let wk = vec![0.01f32; k_dim * hidden];
2541 let wv = vec![0.01f32; k_dim * hidden];
2542 let wo = vec![0.01f32; hidden * q_dim];
2543 add_raw(
2544 "model.layers.0.self_attn.q_proj.weight",
2545 vec![q_dim, hidden],
2546 &wq,
2547 );
2548 add_raw(
2549 "model.layers.0.self_attn.k_proj.weight",
2550 vec![k_dim, hidden],
2551 &wk,
2552 );
2553 add_raw(
2554 "model.layers.0.self_attn.v_proj.weight",
2555 vec![k_dim, hidden],
2556 &wv,
2557 );
2558 add_raw(
2559 "model.layers.0.self_attn.o_proj.weight",
2560 vec![hidden, q_dim],
2561 &wo,
2562 );
2563 let g = vec![0.01f32; inter * hidden];
2564 let d = vec![0.01f32; hidden * inter];
2565 add_raw(
2566 "model.layers.0.mlp.gate_proj.weight",
2567 vec![inter, hidden],
2568 &g,
2569 );
2570 add_raw("model.layers.0.mlp.up_proj.weight", vec![inter, hidden], &g);
2571 add_raw(
2572 "model.layers.0.mlp.down_proj.weight",
2573 vec![hidden, inter],
2574 &d,
2575 );
2576 add_raw("model.norm.weight", vec![hidden], &n1);
2577 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
2578
2579 let cfg = json!({
2580 "format": "aria-quant-bundle",
2581 "format_version": 2,
2582 "quantization": "test",
2583 "group_size_default": 32,
2584 "hadamard_seed": 0,
2585 "model": {
2586 "hidden_size": hidden,
2587 "num_layers": layers,
2588 "num_attention_heads": n_heads,
2589 "num_kv_heads": n_kv,
2590 "intermediate_size": inter,
2591 "vocab_size": vocab,
2592 "context_length": 32,
2593 "rope_theta": 10000.0
2594 },
2595 "tensors": tensors
2596 });
2597 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
2598 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
2599
2600 let mut s = SessionBuilder::new()
2601 .model(dir.path())
2602 .family("qwen/qwen3-0.6b")
2603 .build()
2604 .unwrap();
2605 let gen = s
2606 .generate(
2607 &[1, 2],
2608 &GenerateOpts {
2609 max_tokens: 2,
2610 temperature: 0.0,
2611 },
2612 )
2613 .unwrap();
2614 assert_eq!(gen.tokens.len(), 2);
2615 }
2616
2617 #[test]
2618 fn materialize_accepts_language_model_prefix_and_pre_ffn_norm() {
2619 let dir = tempfile::tempdir().unwrap();
2621 let hidden = 8usize;
2622 let layers = 1usize;
2623 let inter = 16usize;
2624 let vocab = 16usize;
2625 let n_heads = 2usize;
2626 let n_kv = 1usize;
2627 let head_dim = 4usize;
2628 let q_dim = n_heads * head_dim;
2629 let k_dim = n_kv * head_dim;
2630
2631 let mut tensors = serde_json::Map::new();
2632 let mut bin = Vec::new();
2633 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
2634 let offset = bin.len();
2635 for &v in data {
2636 bin.extend_from_slice(&v.to_le_bytes());
2637 }
2638 let nbytes = data.len() * 4;
2639 let mut meta = serde_json::Map::new();
2640 meta.insert("kind".into(), json!("raw"));
2641 meta.insert("dtype".into(), json!("f32"));
2642 meta.insert("shape".into(), json!(shape));
2643 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
2644 tensors.insert(name.to_string(), Value::Object(meta));
2645 };
2646 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
2647 let p = "model.language_model";
2648 add_raw(
2649 &format!("{p}.embed_tokens.weight"),
2650 vec![vocab, hidden],
2651 &emb,
2652 );
2653 let n1 = vec![1.0f32; hidden];
2654 add_raw(
2655 &format!("{p}.layers.0.input_layernorm.weight"),
2656 vec![hidden],
2657 &n1,
2658 );
2659 add_raw(
2660 &format!("{p}.layers.0.pre_feedforward_layernorm.weight"),
2661 vec![hidden],
2662 &n1,
2663 );
2664 let wq = vec![0.01f32; q_dim * hidden];
2665 let wk = vec![0.01f32; k_dim * hidden];
2666 let wv = vec![0.01f32; k_dim * hidden];
2667 let wo = vec![0.01f32; hidden * q_dim];
2668 add_raw(
2669 &format!("{p}.layers.0.self_attn.q_proj.weight"),
2670 vec![q_dim, hidden],
2671 &wq,
2672 );
2673 add_raw(
2674 &format!("{p}.layers.0.self_attn.k_proj.weight"),
2675 vec![k_dim, hidden],
2676 &wk,
2677 );
2678 add_raw(
2679 &format!("{p}.layers.0.self_attn.v_proj.weight"),
2680 vec![k_dim, hidden],
2681 &wv,
2682 );
2683 add_raw(
2684 &format!("{p}.layers.0.self_attn.o_proj.weight"),
2685 vec![hidden, q_dim],
2686 &wo,
2687 );
2688 let g = vec![0.01f32; inter * hidden];
2689 let d = vec![0.01f32; hidden * inter];
2690 add_raw(
2691 &format!("{p}.layers.0.mlp.gate_proj.weight"),
2692 vec![inter, hidden],
2693 &g,
2694 );
2695 add_raw(
2696 &format!("{p}.layers.0.mlp.up_proj.weight"),
2697 vec![inter, hidden],
2698 &g,
2699 );
2700 add_raw(
2701 &format!("{p}.layers.0.mlp.down_proj.weight"),
2702 vec![hidden, inter],
2703 &d,
2704 );
2705 add_raw(&format!("{p}.norm.weight"), vec![hidden], &n1);
2706 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
2707
2708 let cfg = json!({
2709 "format": "aria-quant-bundle",
2710 "format_version": 2,
2711 "quantization": "test",
2712 "group_size_default": 32,
2713 "hadamard_seed": 0,
2714 "model": {
2715 "hidden_size": hidden,
2716 "num_layers": layers,
2717 "num_attention_heads": n_heads,
2718 "num_kv_heads": n_kv,
2719 "intermediate_size": inter,
2720 "vocab_size": vocab,
2721 "context_length": 32,
2722 "rope_theta": 10000.0,
2723 "head_dim": head_dim,
2724 "global_head_dim": head_dim,
2725 "sliding_window": 512,
2726 "partial_rotary_factor": 0.25,
2727 "layer_types": ["full_attention"]
2728 },
2729 "tensors": tensors
2730 });
2731 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
2732 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
2733
2734 let mut s = SessionBuilder::new()
2735 .model(dir.path())
2736 .family("gemma/gemma-4-e2b-it")
2737 .build()
2738 .unwrap();
2739 let gen = s
2740 .generate(
2741 &[1, 2],
2742 &GenerateOpts {
2743 max_tokens: 2,
2744 temperature: 0.0,
2745 },
2746 )
2747 .unwrap();
2748 assert_eq!(gen.tokens.len(), 2);
2749 }
2750
2751 #[test]
2752 fn gemma4_style_double_wide_mlp_and_shared_kv() {
2753 let dir = tempfile::tempdir().unwrap();
2756 let hidden = 8usize;
2757 let layers = 2usize;
2758 let inter = 16usize;
2759 let inter_wide = 32usize;
2760 let vocab = 16usize;
2761 let n_heads = 2usize;
2762 let n_kv = 1usize;
2763 let head_dim = 4usize;
2764 let q_dim = n_heads * head_dim;
2765 let k_dim = n_kv * head_dim;
2766 let p = "model.language_model";
2767
2768 let mut tensors = serde_json::Map::new();
2769 let mut bin = Vec::new();
2770 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
2771 let offset = bin.len();
2772 for &v in data {
2773 bin.extend_from_slice(&v.to_le_bytes());
2774 }
2775 let nbytes = data.len() * 4;
2776 let mut meta = serde_json::Map::new();
2777 meta.insert("kind".into(), json!("raw"));
2778 meta.insert("dtype".into(), json!("f32"));
2779 meta.insert("shape".into(), json!(shape));
2780 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
2781 tensors.insert(name.to_string(), Value::Object(meta));
2782 };
2783 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
2784 add_raw(
2785 &format!("{p}.embed_tokens.weight"),
2786 vec![vocab, hidden],
2787 &emb,
2788 );
2789 let n1 = vec![1.0f32; hidden];
2790 let wq = vec![0.01f32; q_dim * hidden];
2791 let wk = vec![0.01f32; k_dim * hidden];
2792 let wv = vec![0.01f32; k_dim * hidden];
2793 let wo = vec![0.01f32; hidden * q_dim];
2794 for li in 0..layers {
2795 let layer_inter = if li == 0 { inter } else { inter_wide };
2796 add_raw(
2797 &format!("{p}.layers.{li}.input_layernorm.weight"),
2798 vec![hidden],
2799 &n1,
2800 );
2801 add_raw(
2802 &format!("{p}.layers.{li}.pre_feedforward_layernorm.weight"),
2803 vec![hidden],
2804 &n1,
2805 );
2806 add_raw(
2807 &format!("{p}.layers.{li}.self_attn.q_proj.weight"),
2808 vec![q_dim, hidden],
2809 &wq,
2810 );
2811 if li == 0 {
2812 add_raw(
2813 &format!("{p}.layers.{li}.self_attn.k_proj.weight"),
2814 vec![k_dim, hidden],
2815 &wk,
2816 );
2817 add_raw(
2818 &format!("{p}.layers.{li}.self_attn.v_proj.weight"),
2819 vec![k_dim, hidden],
2820 &wv,
2821 );
2822 }
2823 add_raw(
2824 &format!("{p}.layers.{li}.self_attn.o_proj.weight"),
2825 vec![hidden, q_dim],
2826 &wo,
2827 );
2828 let g = vec![0.01f32; layer_inter * hidden];
2829 let d = vec![0.01f32; hidden * layer_inter];
2830 add_raw(
2831 &format!("{p}.layers.{li}.mlp.gate_proj.weight"),
2832 vec![layer_inter, hidden],
2833 &g,
2834 );
2835 add_raw(
2836 &format!("{p}.layers.{li}.mlp.up_proj.weight"),
2837 vec![layer_inter, hidden],
2838 &g,
2839 );
2840 add_raw(
2841 &format!("{p}.layers.{li}.mlp.down_proj.weight"),
2842 vec![hidden, layer_inter],
2843 &d,
2844 );
2845 }
2846 add_raw(&format!("{p}.norm.weight"), vec![hidden], &n1);
2847 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
2848
2849 let cfg = json!({
2850 "format": "aria-quant-bundle",
2851 "format_version": 2,
2852 "quantization": "test",
2853 "group_size_default": 32,
2854 "hadamard_seed": 0,
2855 "model": {
2856 "hidden_size": hidden,
2857 "num_layers": layers,
2858 "num_attention_heads": n_heads,
2859 "num_kv_heads": n_kv,
2860 "intermediate_size": inter,
2861 "vocab_size": vocab,
2862 "context_length": 32,
2863 "rope_theta": 10000.0,
2864 "num_kv_shared_layers": 1,
2865 "head_dim": head_dim,
2866 "global_head_dim": head_dim,
2867 "sliding_window": 512,
2868 "partial_rotary_factor": 0.25,
2869 "layer_types": ["full_attention", "full_attention"]
2870 },
2871 "tensors": tensors
2872 });
2873 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
2874 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
2875
2876 let mut s = SessionBuilder::new()
2877 .model(dir.path())
2878 .family("gemma/gemma-4-e2b-it")
2879 .build()
2880 .unwrap();
2881 let gen = s
2882 .generate(
2883 &[1, 2],
2884 &GenerateOpts {
2885 max_tokens: 2,
2886 temperature: 0.0,
2887 },
2888 )
2889 .unwrap();
2890 assert_eq!(gen.tokens.len(), 2);
2891 }
2892
2893 #[test]
2894 fn stage_b_arch_classes_generate() {
2895 for (path, arch) in arch_class_representatives() {
2896 if matches!(arch, ArchClass::VL | ArchClass::VLA | ArchClass::TextMoE) {
2897 continue; }
2899 if path.contains("qwen3.5") || path.contains("bonsai") {
2901 let dir = tempfile::tempdir().unwrap();
2902 write_tiny_q4_bundle(dir.path()).unwrap();
2903 let err = SessionBuilder::new()
2904 .model(dir.path())
2905 .family(*path)
2906 .build()
2907 .unwrap_err();
2908 assert!(
2909 matches!(err, EngineError::Unsupported(_)),
2910 "{path}: {err:?}"
2911 );
2912 continue;
2913 }
2914 assert!(require_stage_b(path).is_ok(), "{path}");
2915 let dir = tempfile::tempdir().unwrap();
2916 write_tiny_q4_bundle(dir.path()).unwrap();
2917 let mut s = SessionBuilder::new()
2918 .model(dir.path())
2919 .family(*path)
2920 .build()
2921 .unwrap();
2922 assert_eq!(s.arch(), *arch);
2923 assert!(!s.graph_hook_name().is_empty());
2924 let gen = s
2925 .generate(
2926 &s.encode_text("ok"),
2927 &GenerateOpts {
2928 max_tokens: 2,
2929 temperature: 0.0,
2930 },
2931 )
2932 .unwrap();
2933 assert!(!gen.tokens.is_empty(), "{path}");
2934 }
2935 }
2936
2937 #[test]
2938 fn stage_c_vl_vla_hooks() {
2939 let dir = tempfile::tempdir().unwrap();
2940 write_tiny_q4_bundle(dir.path()).unwrap();
2941 let s = SessionBuilder::new()
2942 .model(dir.path())
2943 .family("lfm/lfm2-vl-450m")
2944 .build()
2945 .unwrap();
2946 let rgb = vec![10u8; 3 * 4 * 4];
2947 let err = s.vision_prefix(&rgb, 4, 4).unwrap_err();
2948 assert!(matches!(err, EngineError::Unsupported(_)));
2949
2950 let vla = SessionBuilder::new()
2951 .model(dir.path())
2952 .family("openvla/openvla-7b")
2953 .build()
2954 .unwrap();
2955 let err = vla.predict_action("move", 7).unwrap_err();
2956 assert!(matches!(err, EngineError::Unsupported(_)));
2957 let emb = vla.embed_text("hello").unwrap();
2958 assert_eq!(emb.len(), vla.config().hidden_size);
2959 }
2960
2961 #[test]
2962 fn unknown_family() {
2963 let err = SessionBuilder::new()
2964 .model("/tmp")
2965 .family("no/such-model")
2966 .build()
2967 .unwrap_err();
2968 assert!(matches!(err, EngineError::UnsupportedFamily(_)));
2969 }
2970
2971 #[test]
2972 fn greedy_deterministic() {
2973 let dir = tempfile::tempdir().unwrap();
2974 write_tiny_q4_bundle(dir.path()).unwrap();
2975 let mut s = SessionBuilder::new()
2976 .model(dir.path())
2977 .family("gemma/gemma-4-e2b-it")
2978 .build()
2979 .unwrap();
2980 let prompt = s.encode_text("hi");
2981 let opts = GenerateOpts {
2982 max_tokens: 3,
2983 temperature: 0.0,
2984 };
2985 let a = s.generate(&prompt, &opts).unwrap();
2986 let b = s.generate(&prompt, &opts).unwrap();
2987 assert_eq!(a.tokens, b.tokens);
2988 assert_eq!(a.tokens.len(), 3);
2989 }
2990
2991 #[test]
2992 fn encode_chat_is_longer_than_raw_user_text() {
2993 let dir = tempfile::tempdir().unwrap();
2994 write_tiny_q4_bundle(dir.path()).unwrap();
2995 let s = SessionBuilder::new()
2996 .model(dir.path())
2997 .family("qwen/qwen3-0.6b")
2998 .build()
2999 .unwrap();
3000 let raw = s.encode_text("Hello");
3001 let chat = s.encode_chat(&[ChatTurn::new("user", "Hello")]);
3002 assert!(
3003 chat.len() > raw.len(),
3004 "chat template should wrap the user turn (raw={}, chat={})",
3005 raw.len(),
3006 chat.len()
3007 );
3008 assert!(
3009 (s.config().rope_theta - 1_000_000.0).abs() < 1.0,
3010 "Qwen3 must not keep Llama-default rope_theta=10000, got {}",
3011 s.config().rope_theta
3012 );
3013 }
3014
3015 #[test]
3016 fn incremental_decode_matches_full_recompute() {
3017 let dir = tempfile::tempdir().unwrap();
3018 write_tiny_q4_bundle(dir.path()).unwrap();
3019 let mut s = SessionBuilder::new()
3020 .model(dir.path())
3021 .family("gemma/gemma-4-e2b-it")
3022 .build()
3023 .unwrap();
3024 let prompt = s.encode_text("hi");
3025 let max_tokens = 5usize;
3026
3027 let mut prefix = prompt.clone();
3029 if prefix.is_empty() {
3030 prefix.push(1);
3031 }
3032 let mut full_tokens = Vec::new();
3033 for _ in 0..max_tokens {
3034 let logits = s.forward(&prefix).unwrap();
3035 let next = argmax(&logits);
3036 full_tokens.push(next);
3037 prefix.push(next);
3038 if s.is_stop_id(next) {
3039 full_tokens.pop();
3040 break;
3041 }
3042 }
3043
3044 let incr = s
3045 .generate(
3046 &prompt,
3047 &GenerateOpts {
3048 max_tokens,
3049 temperature: 0.0,
3050 },
3051 )
3052 .unwrap();
3053 assert_eq!(
3054 incr.tokens, full_tokens,
3055 "incremental decode must match full-recompute greedy tokens"
3056 );
3057 }
3058
3059 #[test]
3060 fn profile_records_load_and_generate() {
3061 let dir = tempfile::tempdir().unwrap();
3062 write_tiny_q4_bundle(dir.path()).unwrap();
3063 let mut s = SessionBuilder::new()
3064 .model(dir.path())
3065 .family("gemma/gemma-4-e2b-it")
3066 .compute(ComputePref::Cpu)
3067 .profile(true)
3068 .build()
3069 .unwrap();
3070 assert!(s.compute_label().contains("cpu"));
3071 let load = s.last_profile().expect("load profile");
3072 assert!(!load.ci_fail);
3073 assert!(load.load.materialize_ms >= 0.0);
3074 s.generate(
3075 &s.encode_text("hi"),
3076 &GenerateOpts {
3077 max_tokens: 2,
3078 temperature: 0.0,
3079 },
3080 )
3081 .unwrap();
3082 let p = s.last_profile().expect("generate profile");
3083 let g = p.generate.as_ref().expect("generate timings");
3084 assert!(g.prefill_ms >= 0.0);
3085 assert!(g.decode_ms >= 0.0);
3086 }
3087
3088 #[test]
3089 fn cuda_greedy_matches_cpu_if_available() {
3090 if resolve_compute(ComputePref::Cuda).is_err() {
3091 return;
3092 }
3093 let dir = tempfile::tempdir().unwrap();
3094 write_tiny_q4_bundle(dir.path()).unwrap();
3095 let prompt_text = "hi";
3096 let opts = GenerateOpts {
3097 max_tokens: 4,
3098 temperature: 0.0,
3099 };
3100 let mut cpu = SessionBuilder::new()
3101 .model(dir.path())
3102 .family("gemma/gemma-4-e2b-it")
3103 .compute(ComputePref::Cpu)
3104 .build()
3105 .unwrap();
3106 let mut gpu = SessionBuilder::new()
3107 .model(dir.path())
3108 .family("gemma/gemma-4-e2b-it")
3109 .compute(ComputePref::Cuda)
3110 .build()
3111 .unwrap();
3112 assert!(gpu.compute_label().contains("cuda"));
3113 let prompt = cpu.encode_text(prompt_text);
3114 let a = cpu.generate(&prompt, &opts).unwrap();
3115 let b = gpu.generate(&prompt, &opts).unwrap();
3116 assert_eq!(
3117 a.tokens, b.tokens,
3118 "CUDA greedy tokens must match CPU (tiny bundle)"
3119 );
3120 }
3121
3122 #[test]
3123 fn max_tokens_zero_rejected() {
3124 let dir = tempfile::tempdir().unwrap();
3125 write_tiny_q4_bundle(dir.path()).unwrap();
3126 let mut s = SessionBuilder::new()
3127 .model(dir.path())
3128 .family("gemma/gemma-4-e2b-it")
3129 .build()
3130 .unwrap();
3131 let err = s
3132 .generate(
3133 &s.encode_text("x"),
3134 &GenerateOpts {
3135 max_tokens: 0,
3136 temperature: 0.0,
3137 },
3138 )
3139 .unwrap_err();
3140 assert!(matches!(err, EngineError::InvalidParam(_)));
3141 }
3142
3143 #[test]
3144 fn moe_family_refuses_dense_stub() {
3145 let dir = tempfile::tempdir().unwrap();
3146 write_tiny_q4_bundle(dir.path()).unwrap();
3147 let err = SessionBuilder::new()
3148 .model(dir.path())
3149 .family("lfm/lfm2-8b-a1b")
3150 .build()
3151 .unwrap_err();
3152 assert!(matches!(err, EngineError::Unsupported(_)));
3153 assert_eq!(
3154 lookup_family("lfm/lfm2-8b-a1b").unwrap().arch,
3155 ArchClass::TextMoE
3156 );
3157 assert_eq!(graph_hook(ArchClass::TextMoE), "text_moe_decoder");
3158 }
3159
3160 #[test]
3161 fn geometry_gates_conv_and_experts() {
3162 let dir = tempfile::tempdir().unwrap();
3164 write_tiny_q4_bundle(dir.path()).unwrap();
3165 let cfg_path = dir.path().join("config.json");
3166 let mut cfg: Value =
3167 serde_json::from_str(&std::fs::read_to_string(&cfg_path).unwrap()).unwrap();
3168 cfg["model"]["layer_types"] = json!(["conv", "full_attention"]);
3169 std::fs::write(&cfg_path, cfg.to_string()).unwrap();
3170 let err = SessionBuilder::new()
3171 .model(dir.path())
3172 .family("lfm/lfm2-350m")
3173 .build()
3174 .unwrap_err();
3175 assert!(
3176 matches!(err, EngineError::Format(_)),
3177 "expected missing conv tensors, got {err:?}"
3178 );
3179
3180 let dir2 = tempfile::tempdir().unwrap();
3182 write_tiny_q4_bundle(dir2.path()).unwrap();
3183 let cfg_path = dir2.path().join("config.json");
3184 let mut cfg: Value =
3185 serde_json::from_str(&std::fs::read_to_string(&cfg_path).unwrap()).unwrap();
3186 cfg["model"]["layer_types"] = json!(["linear_attention", "full_attention"]);
3187 std::fs::write(&cfg_path, cfg.to_string()).unwrap();
3188 let err = SessionBuilder::new()
3189 .model(dir2.path())
3190 .family("gemma/gemma-3-270m-it")
3191 .build()
3192 .unwrap_err();
3193 assert!(
3194 matches!(err, EngineError::Format(_)),
3195 "expected missing DeltaNet tensors, got {err:?}"
3196 );
3197 }
3198
3199 #[test]
3200 fn lfm_short_conv_and_attn_generate() {
3201 let dir = tempfile::tempdir().unwrap();
3202 let hidden = 8usize;
3203 let layers = 2usize;
3204 let inter = 16usize;
3205 let vocab = 16usize;
3206 let n_heads = 2usize;
3207 let n_kv = 1usize;
3208 let head_dim = 4usize;
3209 let q_dim = n_heads * head_dim;
3210 let k_dim = n_kv * head_dim;
3211 let kernel = 3usize;
3212
3213 let mut tensors = serde_json::Map::new();
3214 let mut bin = Vec::new();
3215 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
3216 let offset = bin.len();
3217 for &v in data {
3218 bin.extend_from_slice(&v.to_le_bytes());
3219 }
3220 let nbytes = data.len() * 4;
3221 let mut meta = serde_json::Map::new();
3222 meta.insert("kind".into(), json!("raw"));
3223 meta.insert("dtype".into(), json!("f32"));
3224 meta.insert("shape".into(), json!(shape));
3225 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
3226 tensors.insert(name.to_string(), Value::Object(meta));
3227 };
3228 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
3229 add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
3230 let n1 = vec![1.0f32; hidden];
3231 add_raw("model.layers.0.operator_norm.weight", vec![hidden], &n1);
3233 add_raw("model.layers.0.ffn_norm.weight", vec![hidden], &n1);
3234 let in_proj = vec![0.02f32; 3 * hidden * hidden];
3235 let out_proj = vec![0.02f32; hidden * hidden];
3236 let conv_w = vec![0.1f32; hidden * kernel];
3237 add_raw(
3238 "model.layers.0.conv.in_proj.weight",
3239 vec![3 * hidden, hidden],
3240 &in_proj,
3241 );
3242 add_raw(
3243 "model.layers.0.conv.out_proj.weight",
3244 vec![hidden, hidden],
3245 &out_proj,
3246 );
3247 add_raw(
3248 "model.layers.0.conv.conv.weight",
3249 vec![hidden, kernel],
3250 &conv_w,
3251 );
3252 let g = vec![0.02f32; inter * hidden];
3253 let d = vec![0.02f32; hidden * inter];
3254 add_raw(
3255 "model.layers.0.mlp.gate_proj.weight",
3256 vec![inter, hidden],
3257 &g,
3258 );
3259 add_raw("model.layers.0.mlp.up_proj.weight", vec![inter, hidden], &g);
3260 add_raw(
3261 "model.layers.0.mlp.down_proj.weight",
3262 vec![hidden, inter],
3263 &d,
3264 );
3265 add_raw("model.layers.1.operator_norm.weight", vec![hidden], &n1);
3267 add_raw(
3268 "model.layers.1.post_attention_layernorm.weight",
3269 vec![hidden],
3270 &n1,
3271 );
3272 let wq = vec![0.02f32; q_dim * hidden];
3273 let wk = vec![0.02f32; k_dim * hidden];
3274 let wv = vec![0.02f32; k_dim * hidden];
3275 let wo = vec![0.02f32; hidden * q_dim];
3276 add_raw(
3277 "model.layers.1.self_attn.q_proj.weight",
3278 vec![q_dim, hidden],
3279 &wq,
3280 );
3281 add_raw(
3282 "model.layers.1.self_attn.k_proj.weight",
3283 vec![k_dim, hidden],
3284 &wk,
3285 );
3286 add_raw(
3287 "model.layers.1.self_attn.v_proj.weight",
3288 vec![k_dim, hidden],
3289 &wv,
3290 );
3291 add_raw(
3292 "model.layers.1.self_attn.o_proj.weight",
3293 vec![hidden, q_dim],
3294 &wo,
3295 );
3296 add_raw(
3297 "model.layers.1.mlp.gate_proj.weight",
3298 vec![inter, hidden],
3299 &g,
3300 );
3301 add_raw("model.layers.1.mlp.up_proj.weight", vec![inter, hidden], &g);
3302 add_raw(
3303 "model.layers.1.mlp.down_proj.weight",
3304 vec![hidden, inter],
3305 &d,
3306 );
3307 add_raw("model.norm.weight", vec![hidden], &n1);
3308 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
3309
3310 let cfg = json!({
3311 "format": "aria-quant-bundle",
3312 "format_version": 2,
3313 "quantization": "test",
3314 "group_size_default": 32,
3315 "hadamard_seed": 0,
3316 "model": {
3317 "hidden_size": hidden,
3318 "num_layers": layers,
3319 "num_attention_heads": n_heads,
3320 "num_kv_heads": n_kv,
3321 "head_dim": head_dim,
3322 "intermediate_size": inter,
3323 "vocab_size": vocab,
3324 "context_length": 32,
3325 "rope_theta": 10000.0,
3326 "conv_l_cache": kernel,
3327 "layer_types": ["conv", "full_attention"]
3328 },
3329 "tensors": tensors
3330 });
3331 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
3332 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
3333
3334 let mut s = SessionBuilder::new()
3335 .model(dir.path())
3336 .family("lfm/lfm2-350m")
3337 .build()
3338 .unwrap();
3339 let gen = s
3340 .generate(
3341 &[1, 2, 3],
3342 &GenerateOpts {
3343 max_tokens: 2,
3344 temperature: 0.0,
3345 },
3346 )
3347 .unwrap();
3348 assert_eq!(gen.tokens.len(), 2);
3349 }
3350
3351 #[test]
3352 fn moe_topk_experts_generate() {
3353 let dir = tempfile::tempdir().unwrap();
3354 let hidden = 8usize;
3355 let layers = 1usize;
3356 let inter = 16usize;
3357 let vocab = 16usize;
3358 let n_heads = 2usize;
3359 let n_kv = 1usize;
3360 let head_dim = 4usize;
3361 let q_dim = n_heads * head_dim;
3362 let k_dim = n_kv * head_dim;
3363 let n_experts = 4usize;
3364 let top_k = 2usize;
3365
3366 let mut tensors = serde_json::Map::new();
3367 let mut bin = Vec::new();
3368 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
3369 let offset = bin.len();
3370 for &v in data {
3371 bin.extend_from_slice(&v.to_le_bytes());
3372 }
3373 let nbytes = data.len() * 4;
3374 let mut meta = serde_json::Map::new();
3375 meta.insert("kind".into(), json!("raw"));
3376 meta.insert("dtype".into(), json!("f32"));
3377 meta.insert("shape".into(), json!(shape));
3378 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
3379 tensors.insert(name.to_string(), Value::Object(meta));
3380 };
3381 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
3382 add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
3383 let n1 = vec![1.0f32; hidden];
3384 add_raw("model.layers.0.input_layernorm.weight", vec![hidden], &n1);
3385 add_raw(
3386 "model.layers.0.post_attention_layernorm.weight",
3387 vec![hidden],
3388 &n1,
3389 );
3390 let wq = vec![0.02f32; q_dim * hidden];
3391 let wk = vec![0.02f32; k_dim * hidden];
3392 let wv = vec![0.02f32; k_dim * hidden];
3393 let wo = vec![0.02f32; hidden * q_dim];
3394 add_raw(
3395 "model.layers.0.self_attn.q_proj.weight",
3396 vec![q_dim, hidden],
3397 &wq,
3398 );
3399 add_raw(
3400 "model.layers.0.self_attn.k_proj.weight",
3401 vec![k_dim, hidden],
3402 &wk,
3403 );
3404 add_raw(
3405 "model.layers.0.self_attn.v_proj.weight",
3406 vec![k_dim, hidden],
3407 &wv,
3408 );
3409 add_raw(
3410 "model.layers.0.self_attn.o_proj.weight",
3411 vec![hidden, q_dim],
3412 &wo,
3413 );
3414 let router: Vec<f32> = (0..n_experts * hidden)
3415 .map(|i| ((i % n_experts) as f32) * 0.1)
3416 .collect();
3417 add_raw(
3418 "model.layers.0.block_sparse_moe.gate.weight",
3419 vec![n_experts, hidden],
3420 &router,
3421 );
3422 let g = vec![0.02f32; inter * hidden];
3423 let d = vec![0.02f32; hidden * inter];
3424 for e in 0..n_experts {
3425 add_raw(
3426 &format!("model.layers.0.block_sparse_moe.experts.{e}.w1.weight"),
3427 vec![inter, hidden],
3428 &g,
3429 );
3430 add_raw(
3431 &format!("model.layers.0.block_sparse_moe.experts.{e}.w3.weight"),
3432 vec![inter, hidden],
3433 &g,
3434 );
3435 add_raw(
3436 &format!("model.layers.0.block_sparse_moe.experts.{e}.w2.weight"),
3437 vec![hidden, inter],
3438 &d,
3439 );
3440 }
3441 add_raw("model.norm.weight", vec![hidden], &n1);
3442 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
3443
3444 let cfg = json!({
3445 "format": "aria-quant-bundle",
3446 "format_version": 2,
3447 "quantization": "test",
3448 "group_size_default": 32,
3449 "hadamard_seed": 0,
3450 "model": {
3451 "hidden_size": hidden,
3452 "num_layers": layers,
3453 "num_attention_heads": n_heads,
3454 "num_kv_heads": n_kv,
3455 "head_dim": head_dim,
3456 "intermediate_size": inter,
3457 "vocab_size": vocab,
3458 "context_length": 32,
3459 "rope_theta": 10000.0,
3460 "num_experts": n_experts,
3461 "num_experts_per_tok": top_k
3462 },
3463 "tensors": tensors
3464 });
3465 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
3466 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
3467
3468 let mut s = SessionBuilder::new()
3469 .model(dir.path())
3470 .family("inkling/inkling-small")
3471 .build()
3472 .unwrap();
3473 assert_eq!(s.arch(), ArchClass::TextMoE);
3474 assert_eq!(s.graph_hook_name(), "text_moe_decoder");
3475 let gen = s
3476 .generate(
3477 &[1, 2],
3478 &GenerateOpts {
3479 max_tokens: 2,
3480 temperature: 0.0,
3481 },
3482 )
3483 .unwrap();
3484 assert_eq!(gen.tokens.len(), 2);
3485 }
3486
3487 #[test]
3488 fn tiny_q4_codebook_weights_unrotate_on_load() {
3489 let dir = tempfile::tempdir().unwrap();
3490 write_tiny_q4_bundle(dir.path()).unwrap();
3491 let b = load_bundle(dir.path()).unwrap();
3492 let w = b.weight_loaded("blk.0.attn_q.weight").unwrap();
3493 assert!(
3494 w.hdm_seed.is_none(),
3495 "reconstruct_weight path stores original-space W for linear()"
3496 );
3497 let mut s = SessionBuilder::new()
3498 .model(dir.path())
3499 .family("gemma/gemma-4-e2b-it")
3500 .build()
3501 .unwrap();
3502 let gen = s
3503 .generate(
3504 &[1, 2],
3505 &GenerateOpts {
3506 max_tokens: 2,
3507 temperature: 0.0,
3508 },
3509 )
3510 .unwrap();
3511 assert_eq!(gen.tokens.len(), 2);
3512 }
3513
3514 #[test]
3515 fn gemma_hidden_act_geglu_and_qk_norm() {
3516 let dir = tempfile::tempdir().unwrap();
3517 let hidden = 8usize;
3518 let layers = 1usize;
3519 let inter = 16usize;
3520 let vocab = 16usize;
3521 let n_heads = 2usize;
3522 let n_kv = 1usize;
3523 let head_dim = 4usize;
3524 let q_dim = n_heads * head_dim;
3525 let k_dim = n_kv * head_dim;
3526 let p = "model.language_model";
3527
3528 let mut tensors = serde_json::Map::new();
3529 let mut bin = Vec::new();
3530 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
3531 let offset = bin.len();
3532 for &v in data {
3533 bin.extend_from_slice(&v.to_le_bytes());
3534 }
3535 let nbytes = data.len() * 4;
3536 let mut meta = serde_json::Map::new();
3537 meta.insert("kind".into(), json!("raw"));
3538 meta.insert("dtype".into(), json!("f32"));
3539 meta.insert("shape".into(), json!(shape));
3540 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
3541 tensors.insert(name.to_string(), Value::Object(meta));
3542 };
3543 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
3544 add_raw(
3545 &format!("{p}.embed_tokens.weight"),
3546 vec![vocab, hidden],
3547 &emb,
3548 );
3549 let n1 = vec![1.0f32; hidden];
3550 let qn = vec![1.0f32; head_dim];
3551 let kn = vec![1.0f32; head_dim];
3552 let wq = vec![0.01f32; q_dim * hidden];
3553 let wk = vec![0.01f32; k_dim * hidden];
3554 let wv = vec![0.01f32; k_dim * hidden];
3555 let wo = vec![0.01f32; hidden * q_dim];
3556 add_raw(
3557 &format!("{p}.layers.0.input_layernorm.weight"),
3558 vec![hidden],
3559 &n1,
3560 );
3561 add_raw(
3562 &format!("{p}.layers.0.pre_feedforward_layernorm.weight"),
3563 vec![hidden],
3564 &n1,
3565 );
3566 add_raw(
3567 &format!("{p}.layers.0.self_attn.q_proj.weight"),
3568 vec![q_dim, hidden],
3569 &wq,
3570 );
3571 add_raw(
3572 &format!("{p}.layers.0.self_attn.k_proj.weight"),
3573 vec![k_dim, hidden],
3574 &wk,
3575 );
3576 add_raw(
3577 &format!("{p}.layers.0.self_attn.v_proj.weight"),
3578 vec![k_dim, hidden],
3579 &wv,
3580 );
3581 add_raw(
3582 &format!("{p}.layers.0.self_attn.o_proj.weight"),
3583 vec![hidden, q_dim],
3584 &wo,
3585 );
3586 add_raw(
3587 &format!("{p}.layers.0.self_attn.q_norm.weight"),
3588 vec![head_dim],
3589 &qn,
3590 );
3591 add_raw(
3592 &format!("{p}.layers.0.self_attn.k_norm.weight"),
3593 vec![head_dim],
3594 &kn,
3595 );
3596 let g = vec![0.01f32; inter * hidden];
3597 let d = vec![0.01f32; hidden * inter];
3598 add_raw(
3599 &format!("{p}.layers.0.mlp.gate_proj.weight"),
3600 vec![inter, hidden],
3601 &g,
3602 );
3603 add_raw(
3604 &format!("{p}.layers.0.mlp.up_proj.weight"),
3605 vec![inter, hidden],
3606 &g,
3607 );
3608 add_raw(
3609 &format!("{p}.layers.0.mlp.down_proj.weight"),
3610 vec![hidden, inter],
3611 &d,
3612 );
3613 add_raw(&format!("{p}.norm.weight"), vec![hidden], &n1);
3614 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
3615
3616 let cfg = json!({
3617 "format": "aria-quant-bundle",
3618 "format_version": 2,
3619 "quantization": "test",
3620 "group_size_default": 32,
3621 "hadamard_seed": 0,
3622 "model": {
3623 "hidden_size": hidden,
3624 "num_layers": layers,
3625 "num_attention_heads": n_heads,
3626 "num_kv_heads": n_kv,
3627 "head_dim": head_dim,
3628 "global_head_dim": head_dim,
3629 "sliding_window": 512,
3630 "partial_rotary_factor": 0.25,
3631 "intermediate_size": inter,
3632 "vocab_size": vocab,
3633 "context_length": 32,
3634 "rope_theta": 10000.0,
3635 "hidden_act": "gelu_pytorch_tanh",
3636 "layer_types": ["full_attention"]
3637 },
3638 "tensors": tensors
3639 });
3640 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
3641 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
3642
3643 let mut s = SessionBuilder::new()
3644 .model(dir.path())
3645 .family("gemma/gemma-4-e2b-it")
3646 .build()
3647 .unwrap();
3648 assert_eq!(s.config().hidden_act.as_deref(), Some("gelu_pytorch_tanh"));
3649 let gen = s
3650 .generate(
3651 &[1, 2],
3652 &GenerateOpts {
3653 max_tokens: 2,
3654 temperature: 0.0,
3655 },
3656 )
3657 .unwrap();
3658 assert_eq!(gen.tokens.len(), 2);
3659 }
3660
3661 #[test]
3662 fn gated_deltanet_and_full_attn_generate() {
3663 let dir = tempfile::tempdir().unwrap();
3664 let hidden = 8usize;
3665 let inter = 16usize;
3666 let vocab = 16usize;
3667 let n_heads = 2usize;
3668 let n_kv = 1usize;
3669 let head_dim = 4usize;
3670 let q_dim = n_heads * head_dim;
3671 let k_dim = n_kv * head_dim;
3672 let n_lin = 2usize;
3673 let hk = 4usize;
3674 let hv = 4usize;
3675 let key_dim = n_lin * hk;
3676 let value_dim = n_lin * hv;
3677 let conv_k = 4usize;
3678 let conv_dim = key_dim * 2 + value_dim;
3679
3680 let mut tensors = serde_json::Map::new();
3681 let mut bin = Vec::new();
3682 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
3683 let offset = bin.len();
3684 for &v in data {
3685 bin.extend_from_slice(&v.to_le_bytes());
3686 }
3687 let nbytes = data.len() * 4;
3688 let mut meta = serde_json::Map::new();
3689 meta.insert("kind".into(), json!("raw"));
3690 meta.insert("dtype".into(), json!("f32"));
3691 meta.insert("shape".into(), json!(shape));
3692 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
3693 tensors.insert(name.to_string(), Value::Object(meta));
3694 };
3695 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
3696 add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
3697 let n1 = vec![1.0f32; hidden];
3698 add_raw("model.layers.0.input_layernorm.weight", vec![hidden], &n1);
3699 add_raw(
3700 "model.layers.0.post_attention_layernorm.weight",
3701 vec![hidden],
3702 &n1,
3703 );
3704 let qkvz = vec![0.02f32; (2 * key_dim + 2 * value_dim) * hidden];
3705 let ba = vec![0.1f32; 2 * n_lin * hidden];
3706 let conv = vec![0.05f32; conv_dim * conv_k];
3707 let a_log = vec![0.5f32; n_lin];
3708 let dt = vec![1.0f32; n_lin];
3709 let outp = vec![0.02f32; hidden * value_dim];
3710 add_raw(
3711 "model.layers.0.linear_attn.in_proj_qkvz.weight",
3712 vec![2 * key_dim + 2 * value_dim, hidden],
3713 &qkvz,
3714 );
3715 add_raw(
3716 "model.layers.0.linear_attn.in_proj_ba.weight",
3717 vec![2 * n_lin, hidden],
3718 &ba,
3719 );
3720 add_raw(
3721 "model.layers.0.linear_attn.conv1d.weight",
3722 vec![conv_dim, conv_k],
3723 &conv,
3724 );
3725 add_raw("model.layers.0.linear_attn.A_log", vec![n_lin], &a_log);
3726 add_raw("model.layers.0.linear_attn.dt_bias", vec![n_lin], &dt);
3727 add_raw(
3728 "model.layers.0.linear_attn.out_proj.weight",
3729 vec![hidden, value_dim],
3730 &outp,
3731 );
3732 let g = vec![0.02f32; inter * hidden];
3733 let d = vec![0.02f32; hidden * inter];
3734 add_raw(
3735 "model.layers.0.mlp.gate_proj.weight",
3736 vec![inter, hidden],
3737 &g,
3738 );
3739 add_raw("model.layers.0.mlp.up_proj.weight", vec![inter, hidden], &g);
3740 add_raw(
3741 "model.layers.0.mlp.down_proj.weight",
3742 vec![hidden, inter],
3743 &d,
3744 );
3745
3746 add_raw("model.layers.1.input_layernorm.weight", vec![hidden], &n1);
3747 add_raw(
3748 "model.layers.1.post_attention_layernorm.weight",
3749 vec![hidden],
3750 &n1,
3751 );
3752 let wq = vec![0.02f32; q_dim * hidden];
3753 let wk = vec![0.02f32; k_dim * hidden];
3754 let wv = vec![0.02f32; k_dim * hidden];
3755 let wo = vec![0.02f32; hidden * q_dim];
3756 add_raw(
3757 "model.layers.1.self_attn.q_proj.weight",
3758 vec![q_dim, hidden],
3759 &wq,
3760 );
3761 add_raw(
3762 "model.layers.1.self_attn.k_proj.weight",
3763 vec![k_dim, hidden],
3764 &wk,
3765 );
3766 add_raw(
3767 "model.layers.1.self_attn.v_proj.weight",
3768 vec![k_dim, hidden],
3769 &wv,
3770 );
3771 add_raw(
3772 "model.layers.1.self_attn.o_proj.weight",
3773 vec![hidden, q_dim],
3774 &wo,
3775 );
3776 add_raw(
3777 "model.layers.1.mlp.gate_proj.weight",
3778 vec![inter, hidden],
3779 &g,
3780 );
3781 add_raw("model.layers.1.mlp.up_proj.weight", vec![inter, hidden], &g);
3782 add_raw(
3783 "model.layers.1.mlp.down_proj.weight",
3784 vec![hidden, inter],
3785 &d,
3786 );
3787 add_raw("model.norm.weight", vec![hidden], &n1);
3788 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
3789
3790 let cfg = json!({
3791 "format": "aria-quant-bundle",
3792 "format_version": 2,
3793 "quantization": "test",
3794 "hadamard_seed": 0,
3795 "model": {
3796 "hidden_size": hidden,
3797 "num_layers": 2,
3798 "num_attention_heads": n_heads,
3799 "num_kv_heads": n_kv,
3800 "head_dim": head_dim,
3801 "intermediate_size": inter,
3802 "vocab_size": vocab,
3803 "context_length": 32,
3804 "rope_theta": 10000.0,
3805 "layer_types": ["linear_attention", "full_attention"]
3806 },
3807 "tensors": tensors
3808 });
3809 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
3810 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
3811 let mut s = SessionBuilder::new()
3812 .model(dir.path())
3813 .family("qwen/qwen3.5-2b")
3814 .build()
3815 .unwrap();
3816 assert_eq!(s.config().partial_rotary_factor, Some(0.25));
3817 assert!(
3818 (s.config().rope_theta - 10_000_000.0).abs() < 1.0,
3819 "Qwen3.5 Llama-default rope_theta must become 1e7, got {}",
3820 s.config().rope_theta
3821 );
3822 let gen = s
3823 .generate(
3824 &[1, 2, 3],
3825 &GenerateOpts {
3826 max_tokens: 2,
3827 temperature: 0.0,
3828 },
3829 )
3830 .unwrap();
3831 assert_eq!(gen.tokens.len(), 2);
3832 }
3833
3834 #[test]
3835 fn gated_deltanet_split_qwen35_projections_generate() {
3836 let dir = tempfile::tempdir().unwrap();
3837 let hidden = 8usize;
3838 let inter = 16usize;
3839 let vocab = 16usize;
3840 let n_heads = 2usize;
3841 let n_kv = 1usize;
3842 let head_dim = 4usize;
3843 let q_dim = n_heads * head_dim;
3844 let k_dim = n_kv * head_dim;
3845 let n_lin = 2usize;
3846 let hk = 4usize;
3847 let hv = 4usize;
3848 let key_dim = n_lin * hk;
3849 let value_dim = n_lin * hv;
3850 let conv_k = 4usize;
3851 let conv_dim = key_dim * 2 + value_dim;
3852
3853 let mut tensors = serde_json::Map::new();
3854 let mut bin = Vec::new();
3855 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
3856 let offset = bin.len();
3857 for &v in data {
3858 bin.extend_from_slice(&v.to_le_bytes());
3859 }
3860 let nbytes = data.len() * 4;
3861 let mut meta = serde_json::Map::new();
3862 meta.insert("kind".into(), json!("raw"));
3863 meta.insert("dtype".into(), json!("f32"));
3864 meta.insert("shape".into(), json!(shape));
3865 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
3866 tensors.insert(name.to_string(), Value::Object(meta));
3867 };
3868 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
3869 add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
3870 let n1 = vec![1.0f32; hidden];
3871 add_raw("model.layers.0.input_layernorm.weight", vec![hidden], &n1);
3872 add_raw(
3873 "model.layers.0.post_attention_layernorm.weight",
3874 vec![hidden],
3875 &n1,
3876 );
3877 let qkv = vec![0.02f32; (2 * key_dim + value_dim) * hidden];
3878 let z = vec![0.02f32; value_dim * hidden];
3879 let proj_b = vec![0.1f32; n_lin * hidden];
3880 let proj_a = vec![0.1f32; n_lin * hidden];
3881 let conv = vec![0.05f32; conv_dim * conv_k];
3882 let a_log = vec![0.5f32; n_lin];
3883 let dt = vec![1.0f32; n_lin];
3884 let outp = vec![0.02f32; hidden * value_dim];
3885 add_raw(
3886 "model.layers.0.linear_attn.in_proj_qkv.weight",
3887 vec![2 * key_dim + value_dim, hidden],
3888 &qkv,
3889 );
3890 add_raw(
3891 "model.layers.0.linear_attn.in_proj_z.weight",
3892 vec![value_dim, hidden],
3893 &z,
3894 );
3895 add_raw(
3896 "model.layers.0.linear_attn.in_proj_b.weight",
3897 vec![n_lin, hidden],
3898 &proj_b,
3899 );
3900 add_raw(
3901 "model.layers.0.linear_attn.in_proj_a.weight",
3902 vec![n_lin, hidden],
3903 &proj_a,
3904 );
3905 add_raw(
3906 "model.layers.0.linear_attn.conv1d.weight",
3907 vec![conv_dim, conv_k],
3908 &conv,
3909 );
3910 add_raw("model.layers.0.linear_attn.A_log", vec![n_lin], &a_log);
3911 add_raw("model.layers.0.linear_attn.dt_bias", vec![n_lin], &dt);
3912 add_raw(
3913 "model.layers.0.linear_attn.out_proj.weight",
3914 vec![hidden, value_dim],
3915 &outp,
3916 );
3917 let g = vec![0.02f32; inter * hidden];
3918 let d = vec![0.02f32; hidden * inter];
3919 add_raw(
3920 "model.layers.0.mlp.gate_proj.weight",
3921 vec![inter, hidden],
3922 &g,
3923 );
3924 add_raw("model.layers.0.mlp.up_proj.weight", vec![inter, hidden], &g);
3925 add_raw(
3926 "model.layers.0.mlp.down_proj.weight",
3927 vec![hidden, inter],
3928 &d,
3929 );
3930
3931 add_raw("model.layers.1.input_layernorm.weight", vec![hidden], &n1);
3932 add_raw(
3933 "model.layers.1.post_attention_layernorm.weight",
3934 vec![hidden],
3935 &n1,
3936 );
3937 let wq = vec![0.02f32; q_dim * hidden];
3938 let wk = vec![0.02f32; k_dim * hidden];
3939 let wv = vec![0.02f32; k_dim * hidden];
3940 let wo = vec![0.02f32; hidden * q_dim];
3941 add_raw(
3942 "model.layers.1.self_attn.q_proj.weight",
3943 vec![q_dim, hidden],
3944 &wq,
3945 );
3946 add_raw(
3947 "model.layers.1.self_attn.k_proj.weight",
3948 vec![k_dim, hidden],
3949 &wk,
3950 );
3951 add_raw(
3952 "model.layers.1.self_attn.v_proj.weight",
3953 vec![k_dim, hidden],
3954 &wv,
3955 );
3956 add_raw(
3957 "model.layers.1.self_attn.o_proj.weight",
3958 vec![hidden, q_dim],
3959 &wo,
3960 );
3961 add_raw("model.layers.1.mlp.gate_proj.weight", vec![inter, hidden], &g);
3962 add_raw("model.layers.1.mlp.up_proj.weight", vec![inter, hidden], &g);
3963 add_raw(
3964 "model.layers.1.mlp.down_proj.weight",
3965 vec![hidden, inter],
3966 &d,
3967 );
3968 add_raw("model.norm.weight", vec![hidden], &n1);
3969 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
3970
3971 let cfg = json!({
3972 "format": "aria-quant-bundle",
3973 "format_version": 2,
3974 "quantization": "test",
3975 "hadamard_seed": 0,
3976 "model": {
3977 "hidden_size": hidden,
3978 "num_layers": 2,
3979 "num_attention_heads": n_heads,
3980 "num_kv_heads": n_kv,
3981 "head_dim": head_dim,
3982 "intermediate_size": inter,
3983 "vocab_size": vocab,
3984 "context_length": 32,
3985 "rope_theta": 10000.0,
3986 "layer_types": ["linear_attention", "full_attention"]
3987 },
3988 "tensors": tensors
3989 });
3990 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
3991 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
3992 let mut s = SessionBuilder::new()
3993 .model(dir.path())
3994 .family("qwen/qwen3.5-0.8b")
3995 .build()
3996 .unwrap();
3997 let gen = s
3998 .generate(
3999 &[1, 2, 3],
4000 &GenerateOpts {
4001 max_tokens: 2,
4002 temperature: 0.0,
4003 },
4004 )
4005 .unwrap();
4006 assert_eq!(gen.tokens.len(), 2);
4007 }
4008
4009 #[test]
4010 fn vision_and_action_consume_bundle_weights() {
4011 let dir = tempfile::tempdir().unwrap();
4012 write_tiny_q4_bundle(dir.path()).unwrap();
4013 let cfg_path = dir.path().join("config.json");
4014 let mut cfg: Value =
4015 serde_json::from_str(&std::fs::read_to_string(&cfg_path).unwrap()).unwrap();
4016 let hidden = cfg["model"]["hidden_size"].as_u64().unwrap() as usize;
4017 let mut tensors = cfg["tensors"].as_object().cloned().unwrap();
4018 let mut bin = std::fs::read(dir.path().join("weight.bin")).unwrap();
4019 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
4020 let offset = bin.len();
4021 for &v in data {
4022 bin.extend_from_slice(&v.to_le_bytes());
4023 }
4024 let nbytes = data.len() * 4;
4025 tensors.insert(
4026 name.to_string(),
4027 json!({
4028 "kind": "raw",
4029 "dtype": "f32",
4030 "shape": shape,
4031 "offsets": { "data": [offset, nbytes] }
4032 }),
4033 );
4034 };
4035 let vis = vec![0.1f32; hidden * 3];
4036 add_raw("mm_projector.weight", vec![hidden, 3], &vis);
4037 let act_dim = 7usize;
4038 let act = vec![0.05f32; act_dim * hidden];
4039 add_raw("action_head.weight", vec![act_dim, hidden], &act);
4040 cfg["tensors"] = Value::Object(tensors);
4041 std::fs::write(&cfg_path, cfg.to_string()).unwrap();
4042 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
4043
4044 let s = SessionBuilder::new()
4045 .model(dir.path())
4046 .family("lfm/lfm2-vl-450m")
4047 .build()
4048 .unwrap();
4049 let rgb = vec![10u8; 3 * 4 * 4];
4050 let pref = s.vision_prefix(&rgb, 4, 4).unwrap();
4051 assert_eq!(pref.len(), hidden);
4052
4053 let vla = SessionBuilder::new()
4054 .model(dir.path())
4055 .family("openvla/openvla-7b")
4056 .build()
4057 .unwrap();
4058 let a = vla.predict_action("move", act_dim).unwrap();
4059 assert_eq!(a.len(), act_dim);
4060 }
4061
4062 #[test]
4063 fn load_real_hf_named_bundle_if_present() {
4064 let Ok(path) = std::env::var("ARIA_SMOKE_BUNDLE") else {
4066 return;
4067 };
4068 let path = std::path::Path::new(&path);
4069 if !path.join("config.json").is_file() {
4070 return;
4071 }
4072 let family = if path.to_string_lossy().contains("gemma-4") {
4073 "gemma/gemma-4-e2b-it"
4074 } else {
4075 "qwen/qwen3-0.6b"
4076 };
4077 let s = SessionBuilder::new()
4078 .model(path)
4079 .family(family)
4080 .build()
4081 .unwrap_or_else(|e| panic!("{family} bundle should materialize: {e}"));
4082 assert!(s.config().num_layers > 0);
4083 assert!(s.config().hidden_size > 0);
4084 if family.contains("gemma-4") && s.config().hidden_size >= 1024 {
4085 assert!(
4086 s.weights.ple.is_some(),
4087 "real Gemma-4 q4 must load codebook PLE"
4088 );
4089 let hidden = s.config().hidden_size;
4090 let vocab = s.config().vocab_size;
4091 assert!(
4092 s.weights.emb.data.len() >= vocab.saturating_mul(hidden),
4093 "embed table too small for vocab={vocab} hidden={hidden}"
4094 );
4095 }
4096 }
4097
4098 #[test]
4099 fn gemma4_four_norm_ple_and_tied_embed_generate() {
4100 let dir = tempfile::tempdir().unwrap();
4101 let hidden = 8usize;
4102 let layers = 1usize;
4103 let inter = 16usize;
4104 let vocab = 16usize;
4105 let n_heads = 2usize;
4106 let n_kv = 1usize;
4107 let head_dim = 4usize;
4108 let q_dim = n_heads * head_dim;
4109 let k_dim = n_kv * head_dim;
4110 let ple_d = 4usize;
4111 let p = "model.language_model";
4112
4113 let mut tensors = serde_json::Map::new();
4114 let mut bin = Vec::new();
4115 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
4116 let offset = bin.len();
4117 for &v in data {
4118 bin.extend_from_slice(&v.to_le_bytes());
4119 }
4120 let nbytes = data.len() * 4;
4121 let mut meta = serde_json::Map::new();
4122 meta.insert("kind".into(), json!("raw"));
4123 meta.insert("dtype".into(), json!("f32"));
4124 meta.insert("shape".into(), json!(shape));
4125 meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
4126 tensors.insert(name.to_string(), Value::Object(meta));
4127 };
4128 let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
4129 add_raw(
4130 &format!("{p}.embed_tokens.weight"),
4131 vec![vocab, hidden],
4132 &emb,
4133 );
4134 let n1 = vec![1.0f32; hidden];
4135 add_raw(
4136 &format!("{p}.layers.0.input_layernorm.weight"),
4137 vec![hidden],
4138 &n1,
4139 );
4140 add_raw(
4141 &format!("{p}.layers.0.post_attention_layernorm.weight"),
4142 vec![hidden],
4143 &n1,
4144 );
4145 add_raw(
4146 &format!("{p}.layers.0.pre_feedforward_layernorm.weight"),
4147 vec![hidden],
4148 &n1,
4149 );
4150 add_raw(
4151 &format!("{p}.layers.0.post_feedforward_layernorm.weight"),
4152 vec![hidden],
4153 &n1,
4154 );
4155 add_raw(&format!("{p}.layers.0.layer_scalar"), vec![1], &[0.5f32]);
4157 let wq = vec![0.01f32; q_dim * hidden];
4158 let wk = vec![0.01f32; k_dim * hidden];
4159 let wv = vec![0.01f32; k_dim * hidden];
4160 let wo = vec![0.01f32; hidden * q_dim];
4161 add_raw(
4162 &format!("{p}.layers.0.self_attn.q_proj.weight"),
4163 vec![q_dim, hidden],
4164 &wq,
4165 );
4166 add_raw(
4167 &format!("{p}.layers.0.self_attn.k_proj.weight"),
4168 vec![k_dim, hidden],
4169 &wk,
4170 );
4171 add_raw(
4172 &format!("{p}.layers.0.self_attn.v_proj.weight"),
4173 vec![k_dim, hidden],
4174 &wv,
4175 );
4176 add_raw(
4177 &format!("{p}.layers.0.self_attn.o_proj.weight"),
4178 vec![hidden, q_dim],
4179 &wo,
4180 );
4181 let g = vec![0.01f32; inter * hidden];
4182 let d = vec![0.01f32; hidden * inter];
4183 add_raw(
4184 &format!("{p}.layers.0.mlp.gate_proj.weight"),
4185 vec![inter, hidden],
4186 &g,
4187 );
4188 add_raw(
4189 &format!("{p}.layers.0.mlp.up_proj.weight"),
4190 vec![inter, hidden],
4191 &g,
4192 );
4193 add_raw(
4194 &format!("{p}.layers.0.mlp.down_proj.weight"),
4195 vec![hidden, inter],
4196 &d,
4197 );
4198 let packed = layers * ple_d;
4199 let ple_emb = vec![0.02f32; vocab * packed];
4200 add_raw(
4201 &format!("{p}.embed_tokens_per_layer.weight"),
4202 vec![vocab, packed],
4203 &ple_emb,
4204 );
4205 let ple_proj = vec![0.01f32; packed * hidden];
4206 add_raw(
4207 &format!("{p}.per_layer_model_projection.weight"),
4208 vec![packed, hidden],
4209 &ple_proj,
4210 );
4211 let ple_pn = vec![1.0f32; ple_d];
4212 add_raw(
4213 &format!("{p}.per_layer_projection_norm.weight"),
4214 vec![ple_d],
4215 &ple_pn,
4216 );
4217 let ple_gate = vec![0.01f32; ple_d * hidden];
4218 let ple_out = vec![0.01f32; hidden * ple_d];
4219 add_raw(
4220 &format!("{p}.layers.0.per_layer_input_gate.weight"),
4221 vec![ple_d, hidden],
4222 &ple_gate,
4223 );
4224 add_raw(
4225 &format!("{p}.layers.0.per_layer_projection.weight"),
4226 vec![hidden, ple_d],
4227 &ple_out,
4228 );
4229 add_raw(
4230 &format!("{p}.layers.0.post_per_layer_input_norm.weight"),
4231 vec![hidden],
4232 &n1,
4233 );
4234 add_raw(&format!("{p}.norm.weight"), vec![hidden], &n1);
4235 add_raw("lm_head.weight", vec![vocab, hidden], &emb);
4236
4237 let cfg = json!({
4238 "format": "aria-quant-bundle",
4239 "format_version": 2,
4240 "quantization": "test",
4241 "group_size_default": 32,
4242 "hadamard_seed": 0,
4243 "model": {
4244 "hidden_size": hidden,
4245 "num_layers": layers,
4246 "num_attention_heads": n_heads,
4247 "num_kv_heads": n_kv,
4248 "intermediate_size": inter,
4249 "vocab_size": vocab,
4250 "context_length": 32,
4251 "rope_theta": 10000.0,
4252 "hidden_act": "gelu_pytorch_tanh",
4253 "tie_word_embeddings": true,
4254 "head_dim": head_dim,
4255 "global_head_dim": head_dim,
4256 "sliding_window": 512,
4257 "partial_rotary_factor": 0.25,
4258 "layer_types": ["full_attention"]
4259 },
4260 "tensors": tensors
4261 });
4262 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
4263 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
4264
4265 let mut s = SessionBuilder::new()
4266 .model(dir.path())
4267 .family("gemma/gemma-4-e2b-it")
4268 .build()
4269 .unwrap();
4270 assert!((s.embed_scale - (hidden as f32).sqrt()).abs() < 1e-5);
4271 assert!(s.weights.ple.is_some());
4272 assert!((s.weights.layers[0].layer_scalar - 0.5).abs() < 1e-6);
4273 assert!(s.weights.layers[0].post_attn_norm.is_some());
4274 assert!(s.weights.layers[0].post_ffn_norm.is_some());
4275 let prompt = vec![1u32, 2];
4276 let batched = s
4277 .generate(
4278 &prompt,
4279 &GenerateOpts {
4280 max_tokens: 3,
4281 temperature: 0.0,
4282 },
4283 )
4284 .unwrap();
4285 let step = s
4286 .generate(
4287 &prompt,
4288 &GenerateOpts {
4289 max_tokens: 3,
4290 temperature: 0.0,
4291 },
4292 )
4293 .unwrap();
4294 assert_eq!(batched.tokens, step.tokens);
4295 assert_eq!(batched.tokens.len(), 3);
4296 assert_eq!(s.config().sliding_window, Some(512));
4297 }
4298
4299 #[test]
4300 fn gemma4_ple_required_gate() {
4301 assert!(!gemma4_requires_ple("gemma/gemma-4-e2b-it", 64));
4302 assert!(gemma4_requires_ple("gemma/gemma-4-e2b-it", 1024));
4303 assert!(gemma4_requires_ple("gemma/gemma-4-e2b-it", 1536));
4304 assert!(!gemma4_requires_ple("qwen/qwen3-0.6b", 1536));
4305 }
4306
4307 #[test]
4308 fn gemma4_e2b_scale_missing_ple_is_hard_error() {
4309 let dir = tempfile::tempdir().unwrap();
4310 let hidden = 1024usize;
4311 let layers = 1usize;
4312 let vocab = 8usize;
4313 let n_heads = 8usize;
4314 let n_kv = 1usize;
4315 let head_dim = 128usize;
4316 let q_dim = n_heads * head_dim;
4317 let k_dim = n_kv * head_dim;
4318 let inter = 32usize;
4319 let p = "model.language_model";
4320
4321 let mut tensors = serde_json::Map::new();
4322 let mut bin = Vec::new();
4323 let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
4324 let offset = bin.len();
4325 for &v in data {
4326 bin.extend_from_slice(&v.to_le_bytes());
4327 }
4328 let nbytes = data.len() * 4;
4329 tensors.insert(
4330 name.to_string(),
4331 json!({
4332 "kind": "raw",
4333 "dtype": "f32",
4334 "shape": shape,
4335 "offsets": { "data": [offset, nbytes] }
4336 }),
4337 );
4338 };
4339 let emb = vec![0.01f32; vocab * hidden];
4340 add_raw(&format!("{p}.embed_tokens.weight"), vec![vocab, hidden], &emb);
4341 let ones = vec![1.0f32; hidden];
4342 add_raw(
4343 &format!("{p}.layers.0.input_layernorm.weight"),
4344 vec![hidden],
4345 &ones,
4346 );
4347 add_raw(
4348 &format!("{p}.layers.0.pre_feedforward_layernorm.weight"),
4349 vec![hidden],
4350 &ones,
4351 );
4352 add_raw(
4353 &format!("{p}.layers.0.post_attention_layernorm.weight"),
4354 vec![hidden],
4355 &ones,
4356 );
4357 add_raw(
4358 &format!("{p}.layers.0.post_feedforward_layernorm.weight"),
4359 vec![hidden],
4360 &ones,
4361 );
4362 add_raw(&format!("{p}.norm.weight"), vec![hidden], &ones);
4363 let q = vec![0.01f32; q_dim * hidden];
4364 let k = vec![0.01f32; k_dim * hidden];
4365 add_raw(
4366 &format!("{p}.layers.0.self_attn.q_proj.weight"),
4367 vec![q_dim, hidden],
4368 &q,
4369 );
4370 add_raw(
4371 &format!("{p}.layers.0.self_attn.k_proj.weight"),
4372 vec![k_dim, hidden],
4373 &k,
4374 );
4375 add_raw(
4376 &format!("{p}.layers.0.self_attn.v_proj.weight"),
4377 vec![k_dim, hidden],
4378 &k,
4379 );
4380 add_raw(
4381 &format!("{p}.layers.0.self_attn.o_proj.weight"),
4382 vec![hidden, q_dim],
4383 &q,
4384 );
4385 let qn = vec![1.0f32; head_dim];
4386 add_raw(
4387 &format!("{p}.layers.0.self_attn.q_norm.weight"),
4388 vec![head_dim],
4389 &qn,
4390 );
4391 add_raw(
4392 &format!("{p}.layers.0.self_attn.k_norm.weight"),
4393 vec![head_dim],
4394 &qn,
4395 );
4396 let g = vec![0.01f32; inter * hidden];
4397 add_raw(
4398 &format!("{p}.layers.0.mlp.gate_proj.weight"),
4399 vec![inter, hidden],
4400 &g,
4401 );
4402 add_raw(
4403 &format!("{p}.layers.0.mlp.up_proj.weight"),
4404 vec![inter, hidden],
4405 &g,
4406 );
4407 add_raw(
4408 &format!("{p}.layers.0.mlp.down_proj.weight"),
4409 vec![hidden, inter],
4410 &g,
4411 );
4412 let cfg = json!({
4413 "format": "aria-quant-bundle",
4414 "format_version": 2,
4415 "quantization": "test",
4416 "group_size_default": 32,
4417 "hadamard_seed": 0,
4418 "model": {
4419 "hidden_size": hidden,
4420 "num_layers": layers,
4421 "num_attention_heads": n_heads,
4422 "num_kv_heads": n_kv,
4423 "intermediate_size": inter,
4424 "vocab_size": vocab,
4425 "context_length": 32,
4426 "rope_theta": 10000.0,
4427 "hidden_act": "gelu_pytorch_tanh",
4428 "tie_word_embeddings": true,
4429 "head_dim": head_dim,
4430 "global_head_dim": head_dim,
4431 "sliding_window": 512,
4432 "partial_rotary_factor": 0.25,
4433 "num_kv_shared_layers": 0,
4434 "layer_types": ["full_attention"]
4435 },
4436 "tensors": tensors
4437 });
4438 std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
4439 std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
4440
4441 let err = SessionBuilder::new()
4442 .model(dir.path())
4443 .family("gemma/gemma-4-e2b-it")
4444 .build()
4445 .unwrap_err();
4446 let msg = err.to_string();
4447 assert!(
4448 msg.contains("PLE") && msg.contains("embed_tokens_per_layer"),
4449 "{msg}"
4450 );
4451 }
4452
4453 #[test]
4454 fn gemma4_sliding_window_config_and_generate() {
4455 let dir_wide = tempfile::tempdir().unwrap();
4456 write_tiny_q4_bundle(dir_wide.path()).unwrap();
4457 let dir_narrow = tempfile::tempdir().unwrap();
4458 write_tiny_q4_bundle(dir_narrow.path()).unwrap();
4459 let patch = |path: &std::path::Path, window: usize| {
4460 let cfg_path = path.join("config.json");
4461 let raw = std::fs::read_to_string(&cfg_path).unwrap();
4462 let mut cfg: Value = serde_json::from_str(&raw).unwrap();
4463 cfg["model"]["sliding_window"] = json!(window);
4464 cfg["model"]["layer_types"] = json!(["sliding_attention", "sliding_attention"]);
4465 std::fs::write(&cfg_path, cfg.to_string()).unwrap();
4466 };
4467 patch(dir_wide.path(), 512);
4468 patch(dir_narrow.path(), 1);
4469 let wide = SessionBuilder::new()
4470 .model(dir_wide.path())
4471 .family("gemma/gemma-4-e2b-it")
4472 .build()
4473 .unwrap();
4474 let mut narrow = SessionBuilder::new()
4475 .model(dir_narrow.path())
4476 .family("gemma/gemma-4-e2b-it")
4477 .build()
4478 .unwrap();
4479 assert_eq!(wide.config().sliding_window, Some(512));
4480 assert_eq!(narrow.config().sliding_window, Some(1));
4481 assert_eq!(wide.attn_window(AttnKind::Sliding), Some(512));
4482 assert_eq!(narrow.attn_window(AttnKind::Sliding), Some(1));
4483 for layer in &narrow.weights.layers {
4484 if let LayerOp::Attn(attn) = &layer.op {
4485 assert_eq!(attn.kind, AttnKind::Sliding);
4486 }
4487 }
4488 let prompt = vec![1u32, 2, 3, 4];
4489 let gen = narrow
4490 .generate(
4491 &prompt,
4492 &GenerateOpts {
4493 max_tokens: 3,
4494 temperature: 0.0,
4495 },
4496 )
4497 .unwrap();
4498 assert_eq!(gen.tokens.len(), 3);
4499
4500 let mut incr = SessionBuilder::new()
4502 .model(dir_narrow.path())
4503 .family("gemma/gemma-4-e2b-it")
4504 .build()
4505 .unwrap();
4506 let again = incr
4507 .generate(
4508 &prompt,
4509 &GenerateOpts {
4510 max_tokens: 3,
4511 temperature: 0.0,
4512 },
4513 )
4514 .unwrap();
4515 assert_eq!(gen.tokens, again.tokens);
4516 }
4517}