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