Skip to main content

aria_inference/
session.rs

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
84/// Persistent caches for autoregressive decode (KV / short-conv / DeltaNet).
85struct 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    /// Next absolute position (RoPE / seq index).
92    pos: usize,
93}
94
95#[derive(Clone)]
96struct AttnWeights {
97    wq: MatWeight,
98    /// None on Gemma-4 KV-consumer layers (reuse producer cache).
99    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    /// Depthwise kernel `[hidden * kernel]`.
113    kernel: Vec<f32>,
114    kernel_size: usize,
115}
116
117/// Qwen3.5 / Bonsai Gated DeltaNet (linear attention).
118#[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    /// HF `layer_scalar` / JAX `skip_scale`. 1.0 when the bundle omits it.
182    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    /// HF `final_logit_softcapping` (Gemma-4 default 30). None = disabled.
207    final_logit_softcap: Option<f32>,
208    tokenizer: Option<BundleTokenizer>,
209    /// Cleared at the start of each `generate`; reused across decode steps.
210    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        // Hub/q4 artifacts may still omit Gemma-4 geometry; fill from HF architecture
279        // then validate. Bundle values always win when present.
280        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        // Materialize with the same config Session uses for RoPE / head_dim /
286        // layer_types so AttnKind and KV sharing match the forward path.
287        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        // Real Gemma-4 E2B/E4B checkpoints always ship PLE. Tiny/unit fixtures may
307        // omit it; materialize already errors if model-level PLE is partial.
308        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    // Hybrid linear-attn families must declare linear_attention / delta layers.
355    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    // MoE families need explicit expert count (router/experts materialize from that).
374    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
412/// HF Gemma-4 E2B/E4B: repeating 4×sliding + 1×full (full at indices 4,9,...,n-1).
413fn 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
425/// Fill Gemma-4 geometry omitted by published hub q4 (base fields only) or
426/// half-updated bundles. Prefer explicit `config_from_hf` values when present.
427fn 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    // Full E2B/E4B use head_dim=256 / global_head_dim=512; tiny fixtures derive.
441    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
461/// After defaults, Gemma-4 must have a complete geometry (bundle or filled).
462fn 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
524/// When bundle declares distinct global vs local head dims, prefer q_proj geometry.
525fn 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    // LFM MoE uses sigmoid routing; Mixtral/Inkling-style uses softmax.
572    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            // Accept [H,K] or squeezed [H,1,K] (same flat length).
610            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            // Equal k/v dims: qkvz = 2*key + 2*value = 4*key.
667            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                    // Dense FFN on this layer (e.g. LFM2-A1B first layers).
787                    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        // Qwen3-0.6B/1.7B and Gemma-4 tie lm_head to embed; prefer embed even
847        // if a separate lm_head tensor exists (often a worse-quantized copy).
848        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            // HF `post_per_layer_input_norm` is RMSNorm(hidden_size), not ple_d.
926            // A wrong-sized weight would silently chunk-norm and corrupt residuals.
927            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    /// Greedy (temperature<=0) generation from prompt token ids.
1068    /// Prefills the prompt once, then runs one incremental decode step per new token.
1069    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                // Stage A: greedy for determinism; temperature reserved.
1101                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    /// Map token ids → UTF-8 via bundle `tokenizer.json` (byte-level when applicable).
1134    /// Falls back to `<id>` placeholders when no sidecar is present.
1135    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    /// Encode with bundle `tokenizer.json` when present; else naive byte fallback.
1146    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    /// Encode OpenAI-style messages with the family / tokenizer chat template.
1158    pub fn encode_chat(&self, messages: &[ChatTurn]) -> Vec<u32> {
1159        // Prefer the session family when it is gemma-4 so a stale tokenizer hint
1160        // (e.g. gemma-3 `<start_of_turn>`) cannot override `<|turn>` markers.
1161        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    /// Mean-pool token embeddings (stage C `/v1/embeddings`).
1189    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    /// Stage C VL: project RGB via bundle vision weights (no mean-pool stub).
1212    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    /// Stage C VLA: project last-token embedding with bundle action weights.
1271    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    /// Stage C ASR stub bound to session vocab.
1304    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    /// Sliding layers attend to the last `sliding_window` keys (Gemma-4: 512).
1349    /// Full-attention layers keep the causal prefix. KV cache itself is not cropped
1350    /// so shared-KV producers still keep full history.
1351    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                    // Validated by require_gemma4_config; 1.0 means full RoPE.
1364                    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            // HF `_compute_proportional_rope_parameters`: rotate leading
1407            // factor*head_dim/2 pairs; inv_freq denominator is full head_dim.
1408            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    /// Full-sequence forward (allocates fresh caches). Used by tests / parity checks.
1624    #[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                    // RMSNormGated approx: rms(core) * silu(z)
2040                    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
2092/// Optional confidence heuristic for hybrid: mean max-softmax over last logits proxy.
2093pub 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        // Hub q4 still ships base geometry only (no Gemma-4 RoPE / window fields).
2135        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        // Minimal HF-named raw bundle matching Qwen-style paths.
2202        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; // q_dim = 8, k_dim = 4
2210        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        // Gemma-4 / Gemma-3n VL-style HF paths.
2318        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        // Config intermediate_size stays at the narrow width; layer 1 is 2× (KV-shared)
2452        // and omits k/v projections — reuses producer KV cache (num_kv_shared_layers=1).
2453        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; // stage C / MoE gated separately
2596            }
2597            // Hybrid linear_attention families refuse dense Session until DeltaNet lands.
2598            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        // Legacy path: re-run full forward over the growing prefix each step.
2726        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        // layer_types=conv without conv.* weights → Format (not silent dense attn).
2861        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        // linear_attention still hard-gated.
2879        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        // Layer 0: short-conv
2930        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        // Layer 1: full attention
2964        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        // Optional local smoke: ARIA_SMOKE_BUNDLE=/path/to/qwen3-0.6b_q4 or gemma-4-e2b-it_q4
3582        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        // HF skip_scale → layer_scalar (not ones at the published checkpoint).
3673        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        // Prefill window slice must match stepwise decode (same mask).
4018        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}