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