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