Skip to main content

aria_inference/
session.rs

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