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