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