use crate::bundle::{load_bundle, Bundle, LoadedWeight};
use crate::chat::{apply_chat_template, strip_assistant_visible, ChatTurn};
use crate::family::{effective_rope_theta, graph_hook, require_runnable, ArchClass, Family};
use crate::multimodal::asr_transcribe_pcm16le;
use crate::profile::{
elapsed_ms, load_profile_begin, load_profile_set_cuda_upload, load_profile_set_materialize,
load_profile_set_mmap, load_profile_take, EngineProfile, GenerateProfile,
};
use crate::tensor_names::{
action_head_names, altup_projection_names, altup_unembed_names, attn_k_names,
attn_k_norm_names, attn_norm_names, attn_o_names, attn_post_norm_names, attn_q_names,
attn_q_norm_names, attn_v_names, attn_v_norm_names, conv_in_proj_names, conv_kernel_names,
conv_out_proj_names, emb_names, embed_per_layer_names, ffn_down_names, ffn_gate_names,
ffn_norm_names, ffn_post_norm_names, ffn_up_names, layer_altup_correct_scale_names,
layer_altup_correction_coef_names, layer_altup_prediction_coef_names, layer_altup_router_names,
layer_altup_router_norm_names, layer_laurel_left_names, layer_laurel_norm_names,
layer_laurel_right_names, layer_ple_gate_names, layer_ple_post_norm_names,
layer_ple_proj_names, layer_scalar_names, linear_a_log_names, linear_conv1d_names,
linear_dt_bias_names, linear_in_proj_a_names, linear_in_proj_b_names, linear_in_proj_ba_names,
linear_in_proj_qkv_names, linear_in_proj_qkvz_names, linear_in_proj_z_names,
linear_out_norm_names, linear_out_proj_names, moe_expert_down_names, moe_expert_gate_names,
moe_expert_up_names, moe_router_names, output_names, output_norm_names,
per_layer_model_projection_names, per_layer_projection_norm_names, pre_feedforward_norm_names,
vision_proj_names,
};
use crate::tokenizer::{decode_placeholders, encode_naive, BundleTokenizer};
use aria_kernel::{
attention_causal_with_scale, attention_with_scale, gated_delta_step, geglu, gelu_pytorch_tanh,
hdm_linear, kv_sliding_view, linear_cpu, moe_topk_route, resolve_compute, rms_norm,
rms_norm_gemma, rope_half, rope_half_partial, rope_half_proportional, short_conv_step,
silu_vec, softplus, swiglu, ComputeBackend, ComputePref, CudaContext, EngineError,
GatedDeltaStep,
};
use std::cell::RefCell;
use std::collections::HashMap;
use std::path::Path;
use std::sync::Arc;
use std::time::Instant;
#[derive(Debug, Clone)]
pub struct GenerateOpts {
pub max_tokens: usize,
pub temperature: f32,
}
impl Default for GenerateOpts {
fn default() -> Self {
Self {
max_tokens: 16,
temperature: 0.0,
}
}
}
#[derive(Debug, Clone)]
pub struct Generation {
pub tokens: Vec<u32>,
pub text: String,
}
#[derive(Clone)]
struct MatWeight {
data: Arc<Vec<f32>>,
hdm_seed: Option<i64>,
}
impl MatWeight {
fn from_loaded(w: LoadedWeight) -> Self {
Self {
data: Arc::new(w.data),
hdm_seed: w.hdm_seed,
}
}
fn concat_out(a: &Self, b: &Self) -> Self {
let mut data = Vec::with_capacity(a.data.len() + b.data.len());
data.extend_from_slice(&a.data);
data.extend_from_slice(&b.data);
Self {
data: Arc::new(data),
hdm_seed: None,
}
}
}
#[derive(Clone, Copy)]
enum GemmAcct {
Attn,
Ffn,
LmHead,
Other,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
enum AttnKind {
Sliding,
Full,
}
#[derive(Debug, Clone, Copy, PartialEq)]
enum RopeMode {
Full,
Partial(f32),
Proportional(f32),
}
struct DecodeState {
k_caches: Vec<Vec<f32>>,
v_caches: Vec<Vec<f32>>,
last_kv_src: HashMap<AttnKind, usize>,
conv_states: Vec<Option<Vec<f32>>>,
delta_states: Vec<Option<Vec<f32>>>,
pos: usize,
}
#[derive(Clone)]
struct AttnWeights {
wq: MatWeight,
wk: Option<MatWeight>,
wv: Option<MatWeight>,
wo: MatWeight,
q_norm: Option<Vec<f32>>,
k_norm: Option<Vec<f32>>,
v_norm: Option<Vec<f32>>,
kind: AttnKind,
q_gate: bool,
}
#[derive(Clone)]
struct ConvWeights {
in_proj: MatWeight,
out_proj: MatWeight,
kernel: Vec<f32>,
kernel_size: usize,
}
#[derive(Clone)]
struct DeltaWeights {
qkvz: MatWeight,
ba: MatWeight,
conv: Vec<f32>,
conv_k: usize,
out_proj: MatWeight,
out_norm: Vec<f32>,
a_log: Vec<f32>,
dt_bias: Vec<f32>,
n_k_heads: usize,
n_v_heads: usize,
head_k: usize,
head_v: usize,
}
#[derive(Clone)]
enum LayerOp {
Attn(AttnWeights),
Conv(ConvWeights),
Linear(DeltaWeights),
}
#[derive(Clone)]
struct ExpertWeights {
gate: MatWeight,
up: MatWeight,
down: MatWeight,
}
#[derive(Clone)]
enum FfnWeights {
Dense {
gate: MatWeight,
up: MatWeight,
down: MatWeight,
},
MoE {
router: MatWeight,
experts: Vec<ExpertWeights>,
top_k: usize,
use_sigmoid: bool,
},
}
struct LayerPle {
gate: MatWeight,
proj: MatWeight,
post_norm: Vec<f32>,
}
struct PleModel {
embed: Arc<Vec<f32>>,
proj: MatWeight,
proj_norm: Vec<f32>,
d: usize,
}
struct LayerAltUp {
modality_router: MatWeight,
router_norm: Vec<f32>,
prediction_coefs: MatWeight,
correction_coefs: MatWeight,
correct_output_scale: Vec<f32>,
}
struct LayerLaurel {
left: MatWeight,
right: MatWeight,
post_norm: Vec<f32>,
rank: usize,
}
struct LayerWeights {
attn_norm: Vec<f32>,
ffn_norm: Vec<f32>,
post_attn_norm: Option<Vec<f32>>,
post_ffn_norm: Option<Vec<f32>>,
ple: Option<LayerPle>,
altup: Option<LayerAltUp>,
laurel: Option<LayerLaurel>,
activation_sparsity: f32,
layer_scalar: f32,
op: LayerOp,
ffn: FfnWeights,
}
struct ModelWeights {
emb: MatWeight,
layers: Vec<LayerWeights>,
output_norm: Vec<f32>,
output: MatWeight,
vision: Option<MatWeight>,
action: Option<MatWeight>,
ple: Option<PleModel>,
altup_projections: Vec<MatWeight>,
altup_unembed: Vec<MatWeight>,
}
pub struct Session {
family: Family,
bundle: Bundle,
weights: ModelWeights,
conf: crate::bundle::ModelConfig,
use_gemma_norm: bool,
use_gemma4: bool,
use_gemma3n: bool,
use_geglu: bool,
embed_scale: f32,
final_logit_softcap: Option<f32>,
tokenizer: Option<BundleTokenizer>,
decode: Option<DecodeState>,
compute: ComputeBackend,
compute_label: String,
cuda: Option<CudaContext>,
profile_on: bool,
last_profile: Option<EngineProfile>,
gen_acc: RefCell<GenerateProfile>,
}
impl std::fmt::Debug for Session {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Session")
.field("family", &self.family)
.field("model", &self.conf.hidden_size)
.finish()
}
}
pub struct SessionBuilder {
path: Option<std::path::PathBuf>,
family_path: String,
compute: ComputePref,
profile: bool,
}
impl SessionBuilder {
pub fn new() -> Self {
Self {
path: None,
family_path: "gemma/gemma-4-e2b-it".into(),
compute: ComputePref::Auto,
profile: false,
}
}
pub fn model(mut self, path: impl AsRef<Path>) -> Self {
self.path = Some(path.as_ref().to_path_buf());
self
}
pub fn family(mut self, path: impl Into<String>) -> Self {
self.family_path = path.into();
self
}
pub fn compute(mut self, pref: ComputePref) -> Self {
self.compute = pref;
self
}
pub fn profile(mut self, on: bool) -> Self {
self.profile = on;
self
}
pub fn build(self) -> Result<Session, EngineError> {
let family = require_runnable(&self.family_path)?;
let _hook = graph_hook(family.arch);
let path = self
.path
.ok_or_else(|| EngineError::InvalidParam("model path required".into()))?;
let (compute, compute_label) = resolve_compute(self.compute)?;
load_profile_begin(self.profile);
let t_mmap = Instant::now();
let bundle = load_bundle(&path)?;
load_profile_set_mmap(elapsed_ms(t_mmap));
let mut conf = bundle.model.clone();
conf.rope_theta = effective_rope_theta(family.path(), conf.rope_theta);
fill_gemma4_architecture_defaults(&mut conf, family.path());
fill_gemma3_architecture_defaults(&mut conf, family.path());
fill_gemma3n_architecture_defaults(&mut conf, family.path());
fill_qwen35_architecture_defaults(&mut conf, family.path());
require_gemma4_config(&conf, family.path())?;
reject_unsupported_geometry(&conf, family)?;
let tokenizer = BundleTokenizer::try_load(&path)?;
let t_mat = Instant::now();
let weights = materialize_with_config(&bundle, family, &conf)?;
require_gemma4_ple(&weights, &conf, family.path())?;
require_gemma3n_altup(&weights, &conf, family.path())?;
load_profile_set_materialize(elapsed_ms(t_mat));
let mut cuda = None;
if compute == ComputeBackend::Cuda {
let t_up = Instant::now();
let ctx = CudaContext::new()?;
upload_weights(&ctx, &weights)?;
load_profile_set_cuda_upload(elapsed_ms(t_up));
cuda = Some(ctx);
}
let act = conf
.hidden_act
.as_deref()
.unwrap_or("")
.to_ascii_lowercase();
let use_gemma4 = family.path().contains("gemma-4");
let use_gemma3n = is_gemma3n(family.path());
let use_gemma_norm = (family.path().contains("gemma") && !use_gemma4 && !use_gemma3n)
|| family.path().contains("qwen3.5");
let use_geglu = act.contains("gelu") || family.path().contains("gemma");
let embed_scale = if family.path().contains("gemma") {
(conf.hidden_size as f32).sqrt()
} else {
1.0
};
let final_logit_softcap = if use_gemma4 || use_gemma3n {
Some(30.0)
} else {
None
};
let load = load_profile_take();
let last_profile = self.profile.then(|| EngineProfile {
compute: compute_label.clone(),
load,
generate: None,
ci_fail: false,
});
Ok(Session {
family,
bundle,
weights,
conf,
use_gemma_norm,
use_gemma4,
use_gemma3n,
use_geglu,
embed_scale,
final_logit_softcap,
tokenizer,
decode: None,
compute,
compute_label,
cuda,
profile_on: self.profile,
last_profile,
gen_acc: RefCell::new(GenerateProfile::default()),
})
}
}
impl Default for SessionBuilder {
fn default() -> Self {
Self::new()
}
}
fn reject_unsupported_geometry(
conf: &crate::bundle::ModelConfig,
family: Family,
) -> Result<(), EngineError> {
let path = family.path();
if path.contains("qwen3.5") || path.contains("bonsai") {
let has_linear = conf
.layer_types
.as_ref()
.map(|t| {
t.iter().any(|s| {
let s = s.to_ascii_lowercase();
s.contains("linear_attention") || s.contains("delta")
})
})
.unwrap_or(false);
if !has_linear {
return Err(EngineError::Unsupported(format!(
"{path}: requires model.layer_types with Gated DeltaNet / linear_attention \
(dense-only bundles are unsupported until DeltaNet lands)"
)));
}
}
if family.is_moe() && conf.num_experts.unwrap_or(0) == 0 {
return Err(EngineError::Unsupported(format!(
"{path}: MoE family requires model.num_experts > 0 in bundle config"
)));
}
Ok(())
}
fn layer_type_str(conf: &crate::bundle::ModelConfig, layer: usize) -> String {
conf.layer_types
.as_ref()
.and_then(|t| t.get(layer))
.map(|s| s.to_ascii_lowercase())
.unwrap_or_else(|| "full_attention".into())
}
fn layer_is_conv(conf: &crate::bundle::ModelConfig, layer: usize) -> bool {
layer_type_str(conf, layer).contains("conv")
}
fn layer_is_linear(conf: &crate::bundle::ModelConfig, layer: usize) -> bool {
let t = layer_type_str(conf, layer);
t.contains("linear_attention") || t.contains("delta")
}
fn attn_kind(conf: &crate::bundle::ModelConfig, layer: usize) -> AttnKind {
if layer_type_str(conf, layer).contains("sliding") {
AttnKind::Sliding
} else {
AttnKind::Full
}
}
fn is_kv_consumer(conf: &crate::bundle::ModelConfig, layer: usize) -> bool {
let n = conf.num_kv_shared_layers.unwrap_or(0);
n > 0 && layer >= conf.num_layers.saturating_sub(n)
}
fn default_gemma4_layer_types(n: usize) -> Vec<String> {
(0..n)
.map(|i| {
if (i + 1) % 5 == 0 {
"full_attention".into()
} else {
"sliding_attention".into()
}
})
.collect()
}
fn is_gemma3_text(path: &str) -> bool {
path.to_ascii_lowercase().contains("gemma-3-")
}
fn is_gemma3n(path: &str) -> bool {
path.to_ascii_lowercase().contains("gemma-3n")
}
const GEMMA3N_ALTUP_N: usize = 4;
fn altup_router_input_scale(hidden: usize) -> f32 {
if hidden == 0 {
0.0
} else {
1.0 / hidden as f32
}
}
fn gemma3n_default_kv_shared(num_layers: usize) -> usize {
num_layers.saturating_sub(20.min(num_layers))
}
fn default_gemma3_layer_types(n: usize) -> Vec<String> {
(0..n)
.map(|i| {
if (i + 1) % 6 == 0 {
"full_attention".into()
} else {
"sliding_attention".into()
}
})
.collect()
}
fn fill_gemma3_architecture_defaults(conf: &mut crate::bundle::ModelConfig, family_path: &str) {
if !is_gemma3_text(family_path) {
return;
}
if conf.layer_types.as_ref().map(|t| t.len()) != Some(conf.num_layers) {
conf.layer_types = Some(default_gemma3_layer_types(conf.num_layers));
}
if conf.sliding_window.unwrap_or(0) == 0 {
conf.sliding_window = Some(512);
}
if conf
.hidden_act
.as_ref()
.map(|s| s.is_empty())
.unwrap_or(true)
{
conf.hidden_act = Some("gelu_pytorch_tanh".into());
}
}
fn fill_gemma3n_architecture_defaults(conf: &mut crate::bundle::ModelConfig, family_path: &str) {
if !is_gemma3n(family_path) {
return;
}
if conf.layer_types.as_ref().map(|t| t.len()) != Some(conf.num_layers) {
conf.layer_types = Some(default_gemma4_layer_types(conf.num_layers));
}
if conf.sliding_window.unwrap_or(0) == 0 {
conf.sliding_window = Some(512);
}
if conf
.hidden_act
.as_ref()
.map(|s| s.is_empty())
.unwrap_or(true)
{
conf.hidden_act = Some("gelu_pytorch_tanh".into());
}
if conf.hidden_size >= 1024 {
if conf.head_dim.unwrap_or(0) == 0 {
conf.head_dim = Some(256);
}
if conf.num_kv_shared_layers.is_none() {
conf.num_kv_shared_layers = Some(gemma3n_default_kv_shared(conf.num_layers));
}
}
}
fn fill_qwen35_architecture_defaults(conf: &mut crate::bundle::ModelConfig, family_path: &str) {
if !family_path.contains("qwen3.5") {
return;
}
if conf.partial_rotary_factor.is_none() {
conf.partial_rotary_factor = Some(0.25);
}
}
fn fill_gemma4_architecture_defaults(conf: &mut crate::bundle::ModelConfig, family_path: &str) {
if !family_path.contains("gemma-4") {
return;
}
if conf.layer_types.as_ref().map(|t| t.len()) != Some(conf.num_layers) {
conf.layer_types = Some(default_gemma4_layer_types(conf.num_layers));
}
if conf.sliding_window.unwrap_or(0) == 0 {
conf.sliding_window = Some(512);
}
if conf.partial_rotary_factor.is_none() {
conf.partial_rotary_factor = Some(0.25);
}
if conf.hidden_size >= 1024 {
if conf.head_dim.unwrap_or(0) == 0 {
conf.head_dim = Some(256);
}
if conf.global_head_dim.unwrap_or(0) == 0 {
conf.global_head_dim = Some(512);
}
if conf.num_kv_shared_layers.is_none() {
conf.num_kv_shared_layers = Some(20);
}
} else {
if conf.head_dim.unwrap_or(0) == 0 && conf.num_attention_heads > 0 {
conf.head_dim = Some(conf.hidden_size / conf.num_attention_heads);
}
if conf.global_head_dim.unwrap_or(0) == 0 {
conf.global_head_dim = conf.head_dim;
}
}
}
fn require_gemma4_config(
conf: &crate::bundle::ModelConfig,
family_path: &str,
) -> Result<(), EngineError> {
if !family_path.contains("gemma-4") {
return Ok(());
}
let missing = |field: &str| {
EngineError::Unsupported(format!(
"{family_path}: model.{field} required after Gemma-4 architecture fill \
(re-quantize with current model config_from_hf)"
))
};
match &conf.layer_types {
None => return Err(missing("layer_types")),
Some(t) if t.len() != conf.num_layers => {
return Err(EngineError::Unsupported(format!(
"{family_path}: model.layer_types length {} != num_layers {}",
t.len(),
conf.num_layers
)));
}
Some(_) => {}
}
if conf.sliding_window.unwrap_or(0) == 0 {
return Err(missing("sliding_window"));
}
match conf.partial_rotary_factor {
Some(f) if f > 0.0 && f <= 1.0 => {}
_ => return Err(missing("partial_rotary_factor")),
}
if conf.head_dim.unwrap_or(0) == 0 {
return Err(missing("head_dim"));
}
if conf.global_head_dim.unwrap_or(0) == 0 {
return Err(missing("global_head_dim"));
}
Ok(())
}
fn gemma4_requires_ple(family_path: &str, hidden: usize) -> bool {
(family_path.contains("gemma-4") || is_gemma3n(family_path)) && hidden >= 1024
}
fn gemma3n_requires_altup(family_path: &str, hidden: usize) -> bool {
is_gemma3n(family_path) && hidden >= 1024
}
fn require_gemma3n_altup(
weights: &ModelWeights,
conf: &crate::bundle::ModelConfig,
family_path: &str,
) -> Result<(), EngineError> {
if !gemma3n_requires_altup(family_path, conf.hidden_size) {
return Ok(());
}
let n_extra = GEMMA3N_ALTUP_N - 1;
if weights.altup_projections.len() != n_extra || weights.altup_unembed.len() != n_extra {
return Err(EngineError::Format(format!(
"{family_path}: Gemma-3n AltUp projections required \
(altup_projections + altup_unembed_projections); refusing silent no-op"
)));
}
for (i, layer) in weights.layers.iter().enumerate() {
if layer.altup.is_none() || layer.laurel.is_none() {
return Err(EngineError::Format(format!(
"{family_path}: Gemma-3n layer {i} missing AltUp/Laurel weights"
)));
}
}
Ok(())
}
fn require_gemma4_ple(
weights: &ModelWeights,
conf: &crate::bundle::ModelConfig,
family_path: &str,
) -> Result<(), EngineError> {
if !gemma4_requires_ple(family_path, conf.hidden_size) {
return Ok(());
}
if weights.ple.is_none() {
return Err(EngineError::Format(format!(
"{family_path}: codebook PLE required for Gemma-4 / Gemma-3n E2B/E4B \
(embed_tokens_per_layer + per_layer_model_projection + \
per_layer_projection_norm); refusing silent no-op"
)));
}
Ok(())
}
fn resolve_attn_kind(
conf: &crate::bundle::ModelConfig,
layer: usize,
q_dim: usize,
n_heads: usize,
) -> AttnKind {
if n_heads > 0 && q_dim.is_multiple_of(n_heads) {
let head_from_q = q_dim / n_heads;
if let (Some(g), Some(h)) = (
conf.global_head_dim.filter(|d| *d > 0),
conf.head_dim.filter(|d| *d > 0),
) {
if g != h {
if head_from_q == g {
return AttnKind::Full;
}
if head_from_q == h {
return AttnKind::Sliding;
}
}
}
}
attn_kind(conf, layer)
}
fn attn_q_geometry(
wq_len: usize,
wo_len: usize,
hidden: usize,
) -> Result<(usize, usize, bool), EngineError> {
if hidden == 0 || !wq_len.is_multiple_of(hidden) {
return Err(EngineError::ShapeMismatch(format!(
"attn q proj weight not divisible by hidden_size (len={wq_len} hidden={hidden})"
)));
}
if !wo_len.is_multiple_of(hidden) {
return Err(EngineError::ShapeMismatch(format!(
"attn output proj weight not divisible by hidden_size (len={wo_len} hidden={hidden})"
)));
}
let q_out = wq_len / hidden;
let wo_in = wo_len / hidden;
if q_out == wo_in {
Ok((q_out, q_out, false))
} else if q_out == 2 * wo_in {
Ok((q_out, wo_in, true))
} else {
Err(EngineError::ShapeMismatch(format!(
"attn output proj weight shape mismatch (wo_len={wo_len} hidden={hidden} q_out={q_out})"
)))
}
}
fn split_interleaved_q_gate(
mixed: &[f32],
seq: usize,
n_heads: usize,
head_dim: usize,
) -> Result<(Vec<f32>, Vec<f32>), EngineError> {
let q_dim = n_heads.saturating_mul(head_dim);
let packed = q_dim.saturating_mul(2);
if packed == 0 || mixed.len() != seq * packed {
return Err(EngineError::ShapeMismatch(format!(
"gated q_proj out {} != seq*2*q_dim {}*{packed}",
mixed.len(),
seq
)));
}
let mut q = vec![0.0f32; seq * q_dim];
let mut gate = vec![0.0f32; seq * q_dim];
for t in 0..seq {
for h in 0..n_heads {
let src = t * packed + h * (2 * head_dim);
let dst = t * q_dim + h * head_dim;
q[dst..dst + head_dim].copy_from_slice(&mixed[src..src + head_dim]);
gate[dst..dst + head_dim].copy_from_slice(&mixed[src + head_dim..src + 2 * head_dim]);
}
}
Ok((q, gate))
}
fn apply_sigmoid_gate(x: &mut [f32], gate: &[f32]) -> Result<(), EngineError> {
if x.len() != gate.len() {
return Err(EngineError::ShapeMismatch(format!(
"attn output gate len {} != attn out {}",
gate.len(),
x.len()
)));
}
for (v, g) in x.iter_mut().zip(gate) {
*v *= 1.0 / (1.0 + (-*g).exp());
}
Ok(())
}
fn materialize_with_config(
b: &Bundle,
family: Family,
conf: &crate::bundle::ModelConfig,
) -> Result<ModelWeights, EngineError> {
fn any_mat(b: &Bundle, names: &[String]) -> Result<MatWeight, EngineError> {
let refs: Vec<&str> = names.iter().map(String::as_str).collect();
Ok(MatWeight::from_loaded(b.weight_loaded_any(&refs)?))
}
fn try_mat(b: &Bundle, names: &[String]) -> Result<Option<MatWeight>, EngineError> {
match any_mat(b, names) {
Ok(w) => Ok(Some(w)),
Err(EngineError::Format(_)) => Ok(None),
Err(e) => Err(e),
}
}
fn any_vec(b: &Bundle, names: &[String]) -> Result<Vec<f32>, EngineError> {
Ok((*any_mat(b, names)?.data).clone())
}
fn optional_vec(b: &Bundle, names: &[String]) -> Option<Vec<f32>> {
any_vec(b, names).ok()
}
let m = conf;
let hidden = m.hidden_size;
let n_heads = m.num_attention_heads;
let n_experts = m.num_experts.unwrap_or(0);
let top_k = m.num_experts_per_tok.unwrap_or(1).max(1);
let use_sigmoid_router = n_experts > 0 && m.layer_types.is_some();
let mut layers = Vec::with_capacity(m.num_layers);
let mut prev_wk: Option<MatWeight> = None;
let mut prev_wv: Option<MatWeight> = None;
for layer in 0..m.num_layers {
let attn_norm = any_vec(b, &attn_norm_names(layer))?;
let pre_ff = optional_vec(b, &pre_feedforward_norm_names(layer));
let post_attn_norm = if pre_ff.is_some() {
optional_vec(b, &attn_post_norm_names(layer))
} else {
None
};
let post_ffn_norm = optional_vec(b, &ffn_post_norm_names(layer));
let ffn_norm = if let Some(v) = pre_ff {
v
} else {
any_vec(b, &ffn_norm_names(layer))?
};
let op = if layer_is_conv(m, layer) {
let in_proj = any_mat(b, &conv_in_proj_names(layer))?;
let out_proj = any_mat(b, &conv_out_proj_names(layer))?;
let kw = any_mat(b, &conv_kernel_names(layer))?;
let kernel_size = m.conv_l_cache.unwrap_or(3).max(1);
if kw.data.len() % hidden != 0 {
return Err(EngineError::ShapeMismatch(format!(
"layer {layer} conv kernel len {} not divisible by hidden {hidden}",
kw.data.len()
)));
}
let inferred_k = kw.data.len() / hidden;
let kernel_size = if inferred_k > 0 {
inferred_k
} else {
kernel_size
};
if kw.data.len() != hidden * kernel_size {
return Err(EngineError::ShapeMismatch(format!(
"layer {layer} conv kernel len {} != hidden*kernel {hidden}*{kernel_size}",
kw.data.len()
)));
}
let kernel = (*kw.data).clone();
if in_proj.data.len() != 3 * hidden * hidden {
return Err(EngineError::ShapeMismatch(format!(
"layer {layer} conv in_proj len {} != 3*hidden*hidden",
in_proj.data.len()
)));
}
if out_proj.data.len() != hidden * hidden {
return Err(EngineError::ShapeMismatch(format!(
"layer {layer} conv out_proj len {} != hidden*hidden",
out_proj.data.len()
)));
}
LayerOp::Conv(ConvWeights {
in_proj,
out_proj,
kernel,
kernel_size,
})
} else if layer_is_linear(m, layer) {
let qkvz = if let Some(w) = try_mat(b, &linear_in_proj_qkvz_names(layer))? {
w
} else {
let qkv = any_mat(b, &linear_in_proj_qkv_names(layer))?;
let z = any_mat(b, &linear_in_proj_z_names(layer))?;
MatWeight::concat_out(&qkv, &z)
};
let ba = if let Some(w) = try_mat(b, &linear_in_proj_ba_names(layer))? {
w
} else {
let proj_b = any_mat(b, &linear_in_proj_b_names(layer))?;
let proj_a = any_mat(b, &linear_in_proj_a_names(layer))?;
MatWeight::concat_out(&proj_b, &proj_a)
};
let conv_w = any_mat(b, &linear_conv1d_names(layer))?;
let out_proj = any_mat(b, &linear_out_proj_names(layer))?;
let a_log = any_vec(b, &linear_a_log_names(layer))?;
let dt_bias = any_vec(b, &linear_dt_bias_names(layer))?;
let n_v_heads = a_log.len();
if n_v_heads == 0 || dt_bias.len() != n_v_heads {
return Err(EngineError::ShapeMismatch(format!(
"layer {layer} A_log/dt_bias head mismatch"
)));
}
if ba.data.len() % hidden != 0 {
return Err(EngineError::ShapeMismatch(
"linear in_proj_ba not divisible by hidden".into(),
));
}
if ba.data.len() / hidden != 2 * n_v_heads {
return Err(EngineError::ShapeMismatch(format!(
"layer {layer} in_proj_ba out {} != 2*n_v_heads {}",
ba.data.len() / hidden,
2 * n_v_heads
)));
}
if qkvz.data.len() % hidden != 0 {
return Err(EngineError::ShapeMismatch(
"linear in_proj_qkvz not divisible by hidden".into(),
));
}
let qkvz_out = qkvz.data.len() / hidden;
if !qkvz_out.is_multiple_of(4) {
return Err(EngineError::ShapeMismatch(format!(
"layer {layer} qkvz out {qkvz_out} not divisible by 4"
)));
}
let key_dim = qkvz_out / 4;
let value_dim = key_dim;
let n_k_heads = n_v_heads;
if n_k_heads == 0
|| !key_dim.is_multiple_of(n_k_heads)
|| !value_dim.is_multiple_of(n_v_heads)
{
return Err(EngineError::ShapeMismatch(format!(
"layer {layer} cannot infer DeltaNet head dims"
)));
}
let head_k = key_dim / n_k_heads;
let head_v = value_dim / n_v_heads;
let conv_dim = key_dim * 2 + value_dim;
if conv_w.data.len() % conv_dim != 0 {
return Err(EngineError::ShapeMismatch(format!(
"layer {layer} conv1d len {} not divisible by conv_dim {conv_dim}",
conv_w.data.len()
)));
}
let conv_k = conv_w.data.len() / conv_dim;
if out_proj.data.len() != hidden * value_dim {
return Err(EngineError::ShapeMismatch(format!(
"layer {layer} linear out_proj len {} != hidden*value_dim",
out_proj.data.len()
)));
}
let out_norm = optional_vec(b, &linear_out_norm_names(layer))
.unwrap_or_else(|| vec![1.0f32; head_v]);
if out_norm.len() != head_v {
return Err(EngineError::ShapeMismatch(format!(
"layer {layer} linear_attn.norm len {} != head_v {head_v}",
out_norm.len()
)));
}
LayerOp::Linear(DeltaWeights {
qkvz,
ba,
conv: (*conv_w.data).clone(),
conv_k,
out_proj,
out_norm,
a_log,
dt_bias,
n_k_heads,
n_v_heads,
head_k,
head_v,
})
} else {
let consumer = is_kv_consumer(m, layer);
let (wk, wv) = if consumer {
(None, None)
} else {
let wk = match any_mat(b, &attn_k_names(layer)) {
Ok(w) => {
prev_wk = Some(w.clone());
Some(w)
}
Err(e) => Some(prev_wk.clone().ok_or_else(|| {
EngineError::Format(format!(
"missing k_proj for layer {layer} and no prior KV to share ({e})"
))
})?),
};
let wv = match any_mat(b, &attn_v_names(layer)) {
Ok(w) => {
prev_wv = Some(w.clone());
Some(w)
}
Err(e) => Some(prev_wv.clone().ok_or_else(|| {
EngineError::Format(format!(
"missing v_proj for layer {layer} and no prior KV to share ({e})"
))
})?),
};
(wk, wv)
};
let wq = any_mat(b, &attn_q_names(layer))?;
let wo = any_mat(b, &attn_o_names(layer))?;
let (_q_out, q_dim, q_gate) = attn_q_geometry(wq.data.len(), wo.data.len(), hidden)?;
LayerOp::Attn(AttnWeights {
wq,
wk,
wv,
wo,
q_norm: optional_vec(b, &attn_q_norm_names(layer)),
k_norm: optional_vec(b, &attn_k_norm_names(layer)),
v_norm: optional_vec(b, &attn_v_norm_names(layer)),
kind: resolve_attn_kind(m, layer, q_dim, n_heads),
q_gate,
})
};
let ffn = if n_experts > 0 {
match any_mat(b, &moe_router_names(layer)) {
Ok(router) => {
if router.data.len() != n_experts * hidden {
return Err(EngineError::ShapeMismatch(format!(
"layer {layer} MoE router len {} != num_experts*hidden {n_experts}*{hidden}",
router.data.len()
)));
}
let mut experts = Vec::with_capacity(n_experts);
for e in 0..n_experts {
experts.push(ExpertWeights {
gate: any_mat(b, &moe_expert_gate_names(layer, e))?,
up: any_mat(b, &moe_expert_up_names(layer, e))?,
down: any_mat(b, &moe_expert_down_names(layer, e))?,
});
}
FfnWeights::MoE {
router,
experts,
top_k,
use_sigmoid: use_sigmoid_router,
}
}
Err(_) => {
FfnWeights::Dense {
gate: any_mat(b, &ffn_gate_names(layer))?,
up: any_mat(b, &ffn_up_names(layer))?,
down: any_mat(b, &ffn_down_names(layer))?,
}
}
}
} else {
FfnWeights::Dense {
gate: any_mat(b, &ffn_gate_names(layer))?,
up: any_mat(b, &ffn_up_names(layer))?,
down: any_mat(b, &ffn_down_names(layer))?,
}
};
let ple = match (
any_mat(b, &layer_ple_gate_names(layer)),
any_mat(b, &layer_ple_proj_names(layer)),
optional_vec(b, &layer_ple_post_norm_names(layer)),
) {
(Ok(gate), Ok(proj), Some(post_norm)) => Some(LayerPle {
gate,
proj,
post_norm,
}),
_ => None,
};
let altup = match (
try_mat(b, &layer_altup_router_names(layer))?,
optional_vec(b, &layer_altup_router_norm_names(layer)),
try_mat(b, &layer_altup_prediction_coef_names(layer))?,
try_mat(b, &layer_altup_correction_coef_names(layer))?,
optional_vec(b, &layer_altup_correct_scale_names(layer)),
) {
(
Some(modality_router),
Some(router_norm),
Some(prediction_coefs),
Some(correction_coefs),
Some(correct_output_scale),
) => Some(LayerAltUp {
modality_router,
router_norm,
prediction_coefs,
correction_coefs,
correct_output_scale,
}),
(None, None, None, None, None) => None,
_ => {
return Err(EngineError::Format(format!(
"layer {layer}: partial Gemma-3n AltUp tensors (need router, norms, pred/corr coefs, correct_output_scale)"
)));
}
};
let laurel = match (
try_mat(b, &layer_laurel_left_names(layer))?,
try_mat(b, &layer_laurel_right_names(layer))?,
optional_vec(b, &layer_laurel_norm_names(layer)),
) {
(Some(left), Some(right), Some(post_norm)) => {
if hidden == 0 || left.data.len() % hidden != 0 {
return Err(EngineError::ShapeMismatch(format!(
"layer {layer} laurel left not divisible by hidden"
)));
}
let rank = left.data.len() / hidden;
if rank == 0 || right.data.len() != hidden * rank {
return Err(EngineError::ShapeMismatch(format!(
"layer {layer} laurel rank/shape mismatch"
)));
}
Some(LayerLaurel {
left,
right,
post_norm,
rank,
})
}
(None, None, None) => None,
_ => {
return Err(EngineError::Format(format!(
"layer {layer}: partial Gemma-3n Laurel tensors"
)));
}
};
let activation_sparsity = if is_gemma3n(family.path()) && m.num_layers >= 30 && layer < 10 {
0.95
} else {
0.0
};
let layer_scalar = optional_vec(b, &layer_scalar_names(layer))
.and_then(|v| v.into_iter().find(|x| x.is_finite()))
.unwrap_or(1.0);
layers.push(LayerWeights {
attn_norm,
ffn_norm,
post_attn_norm,
post_ffn_norm,
ple,
altup,
laurel,
activation_sparsity,
layer_scalar,
op,
ffn,
});
}
let emb_n = emb_names();
let out_norm_n = output_norm_names();
let out_n = output_names();
let vis_n: Vec<String> = vision_proj_names()
.iter()
.map(|s| (*s).to_string())
.collect();
let act_n: Vec<String> = action_head_names()
.iter()
.map(|s| (*s).to_string())
.collect();
let emb = MatWeight::from_loaded(b.weight_loaded_any(&emb_n)?);
let output = if m.tie_word_embeddings.unwrap_or(false)
|| family.path().contains("gemma-4")
|| is_gemma3n(family.path())
|| (family.path().contains("qwen3") && !family.path().contains("qwen3.5"))
{
emb.clone()
} else {
MatWeight::from_loaded(b.weight_loaded_any(&out_n)?)
};
let require_ple = gemma4_requires_ple(family.path(), hidden);
let ple = {
let embed_n = embed_per_layer_names();
let proj_n: Vec<String> = per_layer_model_projection_names()
.iter()
.map(|s| (*s).to_string())
.collect();
let norm_n: Vec<String> = per_layer_projection_norm_names()
.iter()
.map(|s| (*s).to_string())
.collect();
let embed_res = b.weight_loaded_any(&embed_n);
let proj_res = any_mat(b, &proj_n);
let proj_norm = optional_vec(b, &norm_n);
match (embed_res, proj_res, proj_norm) {
(Ok(embed), Ok(proj), Some(proj_norm)) => {
let d = proj_norm.len();
if d == 0 {
return Err(EngineError::ShapeMismatch(
"PLE projection norm dim is 0".into(),
));
}
Some(PleModel {
embed: Arc::new(embed.data),
proj,
proj_norm,
d,
})
}
(embed_res, proj_res, proj_norm) if require_ple => {
let embed_s = match &embed_res {
Ok(_) => "ok".to_string(),
Err(e) => e.to_string(),
};
let proj_s = match &proj_res {
Ok(_) => "ok".to_string(),
Err(e) => e.to_string(),
};
let norm_s = if proj_norm.is_some() { "ok" } else { "missing" };
return Err(EngineError::Format(format!(
"{}: codebook PLE required (embed_tokens_per_layer={embed_s}, \
per_layer_model_projection={proj_s}, per_layer_projection_norm={norm_s})",
family.path()
)));
}
_ => None,
}
};
if let Some(ple) = &ple {
let packed = m.num_layers.saturating_mul(ple.d);
if packed == 0
|| !ple.embed.len().is_multiple_of(packed)
|| ple.proj.data.len() != packed * hidden
{
return Err(EngineError::ShapeMismatch(format!(
"PLE shapes: embed {} proj {} expected packed={} hidden={hidden}",
ple.embed.len(),
ple.proj.data.len(),
packed
)));
}
for (i, layer) in layers.iter().enumerate() {
let Some(lp) = &layer.ple else {
return Err(EngineError::Format(format!(
"PLE model tensors present but layer {i} missing gate/proj/norm"
)));
};
if lp.gate.data.len() != ple.d * hidden || lp.proj.data.len() != hidden * ple.d {
return Err(EngineError::ShapeMismatch(format!(
"layer {i} PLE gate/proj shape mismatch (d={}, hidden={hidden})",
ple.d
)));
}
if lp.post_norm.len() != hidden {
return Err(EngineError::ShapeMismatch(format!(
"layer {i} PLE post_norm len {} != hidden {hidden}",
lp.post_norm.len()
)));
}
}
}
let n_extra = GEMMA3N_ALTUP_N - 1;
let mut altup_projections = Vec::new();
let mut altup_unembed = Vec::new();
let mut altup_partial = false;
for i in 0..n_extra {
match try_mat(b, &altup_projection_names(i))? {
Some(w) => altup_projections.push(w),
None => altup_partial = true,
}
match try_mat(b, &altup_unembed_names(i))? {
Some(w) => altup_unembed.push(w),
None => altup_partial = true,
}
}
if altup_partial {
if !altup_projections.is_empty() || !altup_unembed.is_empty() {
return Err(EngineError::Format(
"partial Gemma-3n altup_projections / altup_unembed_projections".into(),
));
}
altup_projections.clear();
altup_unembed.clear();
} else {
for w in altup_projections.iter().chain(altup_unembed.iter()) {
if w.data.len() != hidden * hidden {
return Err(EngineError::ShapeMismatch(format!(
"Gemma-3n altup projection len {} != hidden² {hidden}",
w.data.len()
)));
}
}
}
Ok(ModelWeights {
emb,
layers,
output_norm: b.weight_loaded_any(&out_norm_n)?.data,
output,
vision: any_mat(b, &vis_n).ok(),
action: any_mat(b, &act_n).ok(),
ple,
altup_projections,
altup_unembed,
})
}
fn upload_weights(ctx: &CudaContext, w: &ModelWeights) -> Result<(), EngineError> {
ctx.upload(&w.emb.data)?;
ctx.upload(&w.output.data)?;
if let Some(v) = &w.vision {
ctx.upload(&v.data)?;
}
if let Some(a) = &w.action {
ctx.upload(&a.data)?;
}
if let Some(ple) = &w.ple {
ctx.upload(&ple.embed)?;
ctx.upload(&ple.proj.data)?;
}
for wproj in w.altup_projections.iter().chain(w.altup_unembed.iter()) {
ctx.upload(&wproj.data)?;
}
for layer in &w.layers {
match &layer.op {
LayerOp::Attn(attn) => {
ctx.upload(&attn.wq.data)?;
if let Some(wk) = &attn.wk {
ctx.upload(&wk.data)?;
}
if let Some(wv) = &attn.wv {
ctx.upload(&wv.data)?;
}
ctx.upload(&attn.wo.data)?;
}
LayerOp::Conv(c) => {
ctx.upload(&c.in_proj.data)?;
ctx.upload(&c.out_proj.data)?;
}
LayerOp::Linear(d) => {
ctx.upload(&d.qkvz.data)?;
ctx.upload(&d.ba.data)?;
ctx.upload(&d.out_proj.data)?;
}
}
if let Some(ple) = &layer.ple {
ctx.upload(&ple.gate.data)?;
ctx.upload(&ple.proj.data)?;
}
if let Some(altup) = &layer.altup {
ctx.upload(&altup.modality_router.data)?;
ctx.upload(&altup.prediction_coefs.data)?;
ctx.upload(&altup.correction_coefs.data)?;
}
if let Some(laurel) = &layer.laurel {
ctx.upload(&laurel.left.data)?;
ctx.upload(&laurel.right.data)?;
}
match &layer.ffn {
FfnWeights::Dense { gate, up, down } => {
ctx.upload(&gate.data)?;
ctx.upload(&up.data)?;
ctx.upload(&down.data)?;
}
FfnWeights::MoE {
router, experts, ..
} => {
ctx.upload(&router.data)?;
for e in experts {
ctx.upload(&e.gate.data)?;
ctx.upload(&e.up.data)?;
ctx.upload(&e.down.data)?;
}
}
}
}
Ok(())
}
impl Session {
pub fn family(&self) -> Family {
self.family
}
pub fn model_id(&self) -> &str {
self.family.path()
}
pub fn config(&self) -> &crate::bundle::ModelConfig {
&self.conf
}
pub fn bundle(&self) -> &Bundle {
&self.bundle
}
pub fn compute_label(&self) -> &str {
&self.compute_label
}
pub fn last_profile(&self) -> Option<&EngineProfile> {
self.last_profile.as_ref()
}
fn wmm(
&self,
w: &MatWeight,
x: &[f32],
out_f: usize,
in_f: usize,
acct: GemmAcct,
) -> Result<Vec<f32>, EngineError> {
let t0 = Instant::now();
let y = if let Some(seed) = w.hdm_seed {
hdm_linear(x, &w.data, out_f, in_f, Some(seed))?
} else if self.compute == ComputeBackend::Cuda {
let ctx = self.cuda.as_ref().ok_or_else(|| {
EngineError::Unsupported("compute=cuda but CudaContext missing".into())
})?;
ctx.linear(x, &w.data, out_f, in_f)?
} else {
linear_cpu(x, &w.data, out_f, in_f)?
};
if self.profile_on {
let ms = elapsed_ms(t0);
let mut g = self.gen_acc.borrow_mut();
match acct {
GemmAcct::Attn => g.gemm_attn_ms += ms,
GemmAcct::Ffn => g.gemm_ffn_ms += ms,
GemmAcct::LmHead => g.gemm_lm_head_ms += ms,
GemmAcct::Other => {}
}
}
Ok(y)
}
fn can_batch_prefill(&self) -> bool {
self.weights.layers.iter().all(|layer| {
matches!(layer.op, LayerOp::Attn(_)) && matches!(layer.ffn, FfnWeights::Dense { .. })
})
}
pub fn generate(
&mut self,
prompt: &[u32],
opts: &GenerateOpts,
) -> Result<Generation, EngineError> {
if opts.max_tokens == 0 {
return Err(EngineError::InvalidParam("max_tokens must be > 0".into()));
}
let mut tokens: Vec<u32> = prompt.to_vec();
if tokens.is_empty() {
tokens.push(1);
}
self.decode = Some(self.fresh_decode_state());
*self.gen_acc.borrow_mut() = GenerateProfile::default();
let result = (|| {
let t_pre = Instant::now();
let mut logits = if self.can_batch_prefill() && tokens.len() > 1 {
self.forward_prompt(&tokens)?
} else {
let mut last = Vec::new();
for &tok in &tokens {
last = self.forward_step(tok)?;
}
last
};
if self.profile_on {
self.gen_acc.borrow_mut().prefill_ms = elapsed_ms(t_pre);
}
let mut generated = Vec::new();
let t_dec = Instant::now();
for _ in 0..opts.max_tokens {
let next = argmax(&logits);
generated.push(next);
tokens.push(next);
if self.is_stop_id(next) {
generated.pop();
break;
}
logits = self.forward_step(next)?;
}
if self.profile_on {
self.gen_acc.borrow_mut().decode_ms = elapsed_ms(t_dec);
}
let text = self.decode_tokens(&generated);
Ok(Generation {
tokens: generated,
text,
})
})();
if self.profile_on {
let mut p = self.last_profile.take().unwrap_or(EngineProfile {
compute: self.compute_label.clone(),
load: load_profile_take(),
generate: None,
ci_fail: false,
});
p.generate = Some(self.gen_acc.borrow().clone());
self.last_profile = Some(p);
}
self.decode = None;
result
}
pub fn decode_tokens(&self, ids: &[u32]) -> String {
match &self.tokenizer {
Some(tok) => {
let raw = tok.decode_opts(ids, false);
strip_assistant_visible(&raw)
}
None => decode_placeholders(ids),
}
}
pub fn encode_text(&self, text: &str) -> Vec<u32> {
match &self.tokenizer {
Some(tok) => match tok.encode(text) {
Ok(ids) if !ids.is_empty() => ids,
Ok(_) => encode_naive(text, self.conf.vocab_size as u32),
Err(_) => encode_naive(text, self.conf.vocab_size as u32),
},
None => encode_naive(text, self.conf.vocab_size as u32),
}
}
pub fn encode_chat(&self, messages: &[ChatTurn]) -> Vec<u32> {
let family = if self.family.path().contains("gemma-4") {
self.family.path()
} else {
self.tokenizer
.as_ref()
.and_then(|t| t.chat_family_hint())
.unwrap_or(self.family.path())
};
let prompt = apply_chat_template(family, messages);
self.encode_text(&prompt)
}
fn is_stop_id(&self, id: u32) -> bool {
match &self.tokenizer {
Some(t) => t.is_stop(id),
None => id == 0,
}
}
pub fn arch(&self) -> ArchClass {
self.family.arch
}
pub fn graph_hook_name(&self) -> &'static str {
graph_hook(self.family.arch)
}
pub fn embed_text(&self, text: &str) -> Result<Vec<f32>, EngineError> {
let toks = self.encode_text(text);
let hidden = self.conf.hidden_size;
let vocab = self.conf.vocab_size;
let mut acc = vec![0.0f32; hidden];
if toks.is_empty() {
return Ok(acc);
}
for &tok in &toks {
let tid = (tok as usize) % vocab;
let row = &self.weights.emb.data[tid * hidden..(tid + 1) * hidden];
for i in 0..hidden {
acc[i] += row[i];
}
}
let inv = 1.0 / toks.len() as f32;
for v in &mut acc {
*v *= inv;
}
Ok(acc)
}
pub fn vision_prefix(
&self,
rgb: &[u8],
height: usize,
width: usize,
) -> Result<Vec<f32>, EngineError> {
if !matches!(self.family.arch, ArchClass::VL | ArchClass::VLA) {
return Err(EngineError::Unsupported(format!(
"vision_prefix not available for arch {:?}",
self.family.arch
)));
}
let Some(proj) = &self.weights.vision else {
return Err(EngineError::Unsupported(format!(
"{}: no vision projector tensor in bundle",
self.family.path()
)));
};
let hidden = self.conf.hidden_size;
if hidden == 0 || proj.data.len() % hidden != 0 {
return Err(EngineError::ShapeMismatch(
"vision projector not divisible by hidden_size".into(),
));
}
let in_f = proj.data.len() / hidden;
let need = height
.checked_mul(width)
.and_then(|n| n.checked_mul(3))
.ok_or_else(|| EngineError::InvalidParam("vision size overflow".into()))?;
if rgb.len() < need {
return Err(EngineError::ShapeMismatch(format!(
"rgb len {} < {}x{}x3",
rgb.len(),
height,
width
)));
}
let mut feat = vec![0.0f32; in_f];
let pixels = height * width;
if in_f == 3 {
let mut acc = [0.0f32; 3];
for p in 0..pixels {
acc[0] += rgb[p * 3] as f32 / 255.0;
acc[1] += rgb[p * 3 + 1] as f32 / 255.0;
acc[2] += rgb[p * 3 + 2] as f32 / 255.0;
}
let s = 1.0 / pixels.max(1) as f32;
feat[0] = acc[0] * s;
feat[1] = acc[1] * s;
feat[2] = acc[2] * s;
} else {
for i in 0..in_f {
feat[i] = rgb[i % need] as f32 / 255.0;
}
}
self.wmm(proj, &feat, hidden, in_f, GemmAcct::Other)
}
pub fn predict_action(&self, prompt: &str, action_dim: usize) -> Result<Vec<f32>, EngineError> {
if self.family.arch != ArchClass::VLA {
return Err(EngineError::Unsupported(format!(
"predict_action requires VLA, got {:?}",
self.family.arch
)));
}
if action_dim == 0 {
return Err(EngineError::InvalidParam("action_dim must be > 0".into()));
}
let Some(head) = &self.weights.action else {
return Err(EngineError::Unsupported(format!(
"{}: no action head tensor in bundle",
self.family.path()
)));
};
let h = self.embed_text(prompt)?;
let hidden = self.conf.hidden_size;
if head.data.len() % hidden != 0 {
return Err(EngineError::ShapeMismatch(
"action head not divisible by hidden_size".into(),
));
}
let out_f = head.data.len() / hidden;
if out_f != action_dim {
return Err(EngineError::ShapeMismatch(format!(
"action head out {out_f} != requested {action_dim}"
)));
}
self.wmm(head, &h, out_f, hidden, GemmAcct::Other)
}
pub fn transcribe_pcm16le(&self, pcm: &[u8]) -> Result<String, EngineError> {
asr_transcribe_pcm16le(pcm, self.conf.vocab_size as u32)
}
fn norm(&self, x: &[f32], weight: &[f32]) -> Result<Vec<f32>, EngineError> {
if self.use_gemma_norm {
rms_norm_gemma(x, weight, 1e-6)
} else {
rms_norm(x, weight, 1e-6)
}
}
fn add_normed_residual(
&self,
x: &mut [f32],
y: &[f32],
post_norm: Option<&[f32]>,
) -> Result<(), EngineError> {
if y.len() != x.len() {
return Err(EngineError::ShapeMismatch(
"residual length mismatch".into(),
));
}
if let Some(w) = post_norm {
let yn = self.norm(y, w)?;
for (a, b) in x.iter_mut().zip(yn.iter()) {
*a += *b;
}
} else {
for (a, b) in x.iter_mut().zip(y.iter()) {
*a += *b;
}
}
Ok(())
}
fn attn_scale(&self, head_dim: usize) -> f32 {
if self.use_gemma4 || self.use_gemma3n {
1.0
} else {
1.0 / (head_dim as f32).sqrt()
}
}
fn attn_window(&self, kind: AttnKind) -> Option<usize> {
if kind != AttnKind::Sliding {
return None;
}
self.conf.sliding_window.filter(|w| *w > 0)
}
fn layer_rope_params(&self, kind: AttnKind) -> (f32, RopeMode) {
if self.use_gemma4 {
match kind {
AttnKind::Sliding => (10_000.0, RopeMode::Full),
AttnKind::Full => {
let factor = self.conf.partial_rotary_factor.unwrap_or(1.0);
let mode = if factor > 0.0 && factor < 1.0 {
RopeMode::Proportional(factor)
} else {
RopeMode::Full
};
(1_000_000.0, mode)
}
}
} else if is_gemma3_text(self.family.path()) || self.use_gemma3n {
match kind {
AttnKind::Sliding => (10_000.0, RopeMode::Full),
AttnKind::Full => {
let theta = if (self.conf.rope_theta - 10_000.0).abs() < 0.5 {
1_000_000.0
} else {
self.conf.rope_theta
};
(theta, RopeMode::Full)
}
}
} else if self.family.path().contains("qwen3.5") {
let factor = self.conf.partial_rotary_factor.unwrap_or(0.25);
let mode = if factor > 0.0 && factor < 1.0 {
RopeMode::Partial(factor)
} else {
RopeMode::Full
};
(self.conf.rope_theta, mode)
} else {
(self.conf.rope_theta, RopeMode::Full)
}
}
fn layer_head_dim(
&self,
kind: AttnKind,
q_dim: usize,
n_heads: usize,
) -> Result<usize, EngineError> {
if n_heads == 0 || !q_dim.is_multiple_of(n_heads) {
return Err(EngineError::ShapeMismatch(
"q_dim not divisible by num_attention_heads".into(),
));
}
let configured = match kind {
AttnKind::Full => self.conf.global_head_dim.or(self.conf.head_dim),
AttnKind::Sliding => self.conf.head_dim,
};
Ok(configured
.filter(|d| *d > 0 && q_dim == n_heads * *d)
.unwrap_or(q_dim / n_heads))
}
fn apply_rope(
x: &mut [f32],
head_dim: usize,
pos: usize,
theta: f32,
mode: RopeMode,
) -> Result<(), EngineError> {
match mode {
RopeMode::Full => rope_half(x, head_dim, pos, theta),
RopeMode::Proportional(factor) => {
rope_half_proportional(x, head_dim, factor, pos, theta)
}
RopeMode::Partial(factor) => {
let rotary_dim = (factor * head_dim as f32) as usize & !1;
if rotary_dim < 2 || rotary_dim >= head_dim {
rope_half(x, head_dim, pos, theta)
} else {
rope_half_partial(x, head_dim, rotary_dim, pos, theta)
}
}
}
}
fn apply_v_norm(
&self,
v: Vec<f32>,
v_norm: Option<&[f32]>,
head_dim: usize,
) -> Result<Vec<f32>, EngineError> {
if let Some(vn) = v_norm {
if vn.len() != head_dim {
return Err(EngineError::ShapeMismatch(format!(
"v_norm len {} != head_dim {head_dim}",
vn.len()
)));
}
self.norm(&v, vn)
} else if self.use_gemma4 || self.use_gemma3n {
let ones = vec![1.0f32; head_dim];
rms_norm(&v, &ones, 1e-6)
} else {
Ok(v)
}
}
fn has_gemma3n_graph(&self) -> bool {
self.use_gemma3n
&& self.weights.altup_projections.len() == GEMMA3N_ALTUP_N - 1
&& self.weights.altup_unembed.len() == GEMMA3N_ALTUP_N - 1
}
fn match_token_magnitude(
src: &[f32],
target: &[f32],
hidden: usize,
) -> Result<Vec<f32>, EngineError> {
if src.len() != target.len() || hidden == 0 || !src.len().is_multiple_of(hidden) {
return Err(EngineError::ShapeMismatch(
"Gemma-3n magnitude match length mismatch".into(),
));
}
let seq = src.len() / hidden;
let mut out = src.to_vec();
for t in 0..seq {
let tb = t * hidden;
let mut t_ms = 0.0f32;
let mut s_ms = 0.0f32;
for i in 0..hidden {
t_ms += target[tb + i] * target[tb + i];
s_ms += src[tb + i] * src[tb + i];
}
let target_mag = (t_ms / hidden as f32).sqrt();
let src_mag = (s_ms / hidden as f32).max(1e-5).sqrt();
let scale = target_mag / src_mag;
for i in 0..hidden {
out[tb + i] *= scale;
}
}
Ok(out)
}
fn gemma3n_expand_streams(
&self,
x0: &[f32],
hidden: usize,
) -> Result<Vec<Vec<f32>>, EngineError> {
let mut streams = vec![x0.to_vec()];
for proj in &self.weights.altup_projections {
let p = self.wmm(proj, x0, hidden, hidden, GemmAcct::Other)?;
streams.push(Self::match_token_magnitude(&p, x0, hidden)?);
}
Ok(streams)
}
fn gemma3n_unembed_streams(
&self,
streams: &[Vec<f32>],
hidden: usize,
) -> Result<Vec<f32>, EngineError> {
if streams.len() != GEMMA3N_ALTUP_N {
return Err(EngineError::Format(
"Gemma-3n unembed expects 4 AltUp streams".into(),
));
}
let mut acc = streams[0].clone();
for (i, proj) in self.weights.altup_unembed.iter().enumerate() {
let u = self.wmm(proj, &streams[i + 1], hidden, hidden, GemmAcct::Other)?;
let matched = Self::match_token_magnitude(&u, &streams[0], hidden)?;
for (a, b) in acc.iter_mut().zip(matched.iter()) {
*a += *b;
}
}
let n = GEMMA3N_ALTUP_N as f32;
for v in &mut acc {
*v /= n;
}
Ok(acc)
}
fn altup_modalities(
&self,
altup: &LayerAltUp,
x: &[f32],
hidden: usize,
) -> Result<Vec<f32>, EngineError> {
if altup.router_norm.len() != hidden {
return Err(EngineError::ShapeMismatch(
"Gemma-3n altup.router_norm dim mismatch".into(),
));
}
let mut xn = self.norm(x, &altup.router_norm)?;
let scale = altup_router_input_scale(hidden);
for v in &mut xn {
*v *= scale;
}
let mut routed = self.wmm(
&altup.modality_router,
&xn,
GEMMA3N_ALTUP_N,
hidden,
GemmAcct::Other,
)?;
for v in &mut routed {
*v = v.tanh();
}
Ok(routed)
}
fn altup_predict(
&self,
layer: &LayerWeights,
streams: &[Vec<f32>],
hidden: usize,
) -> Result<Vec<Vec<f32>>, EngineError> {
let altup = layer
.altup
.as_ref()
.ok_or_else(|| EngineError::Format("Gemma-3n layer missing AltUp".into()))?;
let n = GEMMA3N_ALTUP_N;
if streams.len() != n {
return Err(EngineError::Format(
"Gemma-3n predict expects 4 streams".into(),
));
}
let seq = streams[0].len() / hidden;
let routed = self.altup_modalities(altup, &streams[0], hidden)?;
let coefs = self.wmm(&altup.prediction_coefs, &routed, n * n, n, GemmAcct::Other)?;
let mut preds = vec![vec![0.0f32; seq * hidden]; n];
for t in 0..seq {
for out_s in 0..n {
for h in 0..hidden {
let mut acc = streams[out_s][t * hidden + h];
for in_s in 0..n {
acc += streams[in_s][t * hidden + h] * coefs[t * n * n + out_s * n + in_s];
}
preds[out_s][t * hidden + h] = acc;
}
}
}
Ok(preds)
}
fn altup_correct(
&self,
layer: &LayerWeights,
predictions: &[Vec<f32>],
activated: &[f32],
hidden: usize,
) -> Result<Vec<Vec<f32>>, EngineError> {
let altup = layer
.altup
.as_ref()
.ok_or_else(|| EngineError::Format("Gemma-3n layer missing AltUp".into()))?;
let n = GEMMA3N_ALTUP_N;
let seq = activated.len() / hidden;
let routed = self.altup_modalities(altup, activated, hidden)?;
let mut coefs = self.wmm(&altup.correction_coefs, &routed, n, n, GemmAcct::Other)?;
for v in &mut coefs {
*v += 1.0;
}
let mut out = Vec::with_capacity(n);
for (si, pred) in predictions.iter().enumerate() {
let mut row = pred.clone();
for t in 0..seq {
let c = coefs[t * n + si];
for h in 0..hidden {
let idx = t * hidden + h;
let innov = activated[idx] - predictions[0][idx];
row[idx] += innov * c;
}
}
out.push(row);
}
Ok(out)
}
fn apply_laurel(
&self,
layer: &LayerWeights,
xn: &[f32],
hidden: usize,
) -> Result<Vec<f32>, EngineError> {
let laurel = layer
.laurel
.as_ref()
.ok_or_else(|| EngineError::Format("Gemma-3n layer missing Laurel".into()))?;
let left = self.wmm(&laurel.left, xn, laurel.rank, hidden, GemmAcct::Ffn)?;
let right = self.wmm(&laurel.right, &left, hidden, laurel.rank, GemmAcct::Ffn)?;
let nrm = self.norm(&right, &laurel.post_norm)?;
let mut out = xn.to_vec();
for (a, b) in out.iter_mut().zip(nrm.iter()) {
*a += *b;
}
Ok(out)
}
fn scale_altup_active(
&self,
layer: &LayerWeights,
x: &mut [f32],
hidden: usize,
) -> Result<(), EngineError> {
let Some(altup) = &layer.altup else {
return Ok(());
};
if altup.correct_output_scale.len() == 1 {
let s = altup.correct_output_scale[0];
for v in x.iter_mut() {
*v *= s;
}
return Ok(());
}
if altup.correct_output_scale.len() != hidden {
return Err(EngineError::ShapeMismatch(format!(
"altup.correct_output_scale len {} != hidden {hidden}",
altup.correct_output_scale.len()
)));
}
let seq = x.len() / hidden;
for t in 0..seq {
for h in 0..hidden {
x[t * hidden + h] *= altup.correct_output_scale[h];
}
}
Ok(())
}
fn gemma3n_after_attn(
&self,
streams: &mut [Vec<f32>],
predictions: &[Vec<f32>],
laurel: &[f32],
ao: &[f32],
layer: &LayerWeights,
li: usize,
ple_tok: Option<&[f32]>,
hidden: usize,
) -> Result<(), EngineError> {
let mut active = predictions[0].clone();
self.add_normed_residual(&mut active, ao, layer.post_attn_norm.as_deref())?;
let inv_sqrt2 = std::f32::consts::FRAC_1_SQRT_2;
let mut mix = vec![0.0f32; active.len()];
for i in 0..active.len() {
mix[i] = (active[i] + laurel[i]) * inv_sqrt2;
}
let xn2 = self.norm(&mix, &layer.ffn_norm)?;
let down = self.apply_ffn(layer, &xn2, hidden)?;
self.add_normed_residual(&mut mix, &down, layer.post_ffn_norm.as_deref())?;
let corrected = self.altup_correct(layer, predictions, &mix, hidden)?;
for (dst, src) in streams.iter_mut().zip(corrected.into_iter()) {
*dst = src;
}
let mut first = streams[0].clone();
self.scale_altup_active(layer, &mut first, hidden)?;
if let Some(delta) = self.ple_delta(&first, layer, li, ple_tok, hidden)? {
for s in streams.iter_mut().skip(1) {
for (a, b) in s.iter_mut().zip(delta.iter()) {
*a += *b;
}
}
}
Ok(())
}
fn compute_ple_inputs(
&self,
toks: &[u32],
embeds: &[f32],
) -> Result<Option<Vec<f32>>, EngineError> {
let Some(ple) = &self.weights.ple else {
return Ok(None);
};
let hidden = self.conf.hidden_size;
let n_layers = self.weights.layers.len();
let d = ple.d;
let packed = n_layers * d;
let seq = toks.len();
if seq == 0 || embeds.len() != seq * hidden {
return Err(EngineError::ShapeMismatch(
"PLE embed sequence length mismatch".into(),
));
}
let scale_lookup = (d as f32).sqrt();
let ple_vocab = ple.embed.len() / packed;
if ple_vocab == 0 {
return Err(EngineError::ShapeMismatch("PLE embed vocab is 0".into()));
}
let mut lookup = vec![0.0f32; seq * packed];
for (t, &tok) in toks.iter().enumerate() {
let tid = (tok as usize) % ple_vocab;
let row = &ple.embed[tid * packed..(tid + 1) * packed];
for i in 0..packed {
lookup[t * packed + i] = row[i] * scale_lookup;
}
}
let proj_scale = (hidden as f32).sqrt().recip();
let mut proj = self.wmm(&ple.proj, embeds, packed, hidden, GemmAcct::Other)?;
for v in &mut proj {
*v *= proj_scale;
}
proj = rms_norm(&proj, &ple.proj_norm, 1e-6)?;
let inv_sqrt2 = std::f32::consts::FRAC_1_SQRT_2;
for i in 0..proj.len() {
proj[i] = (proj[i] + lookup[i]) * inv_sqrt2;
}
Ok(Some(proj))
}
fn ple_delta(
&self,
x: &[f32],
layer: &LayerWeights,
li: usize,
ple_tok: Option<&[f32]>,
hidden: usize,
) -> Result<Option<Vec<f32>>, EngineError> {
let (Some(ple), Some(ple_tok)) = (&layer.ple, ple_tok) else {
return Ok(None);
};
let d = self
.weights
.ple
.as_ref()
.map(|p| p.d)
.ok_or_else(|| EngineError::Format("layer PLE without model PLE".into()))?;
let n_layers = self.weights.layers.len();
let seq = x.len() / hidden;
let gate_out = self.wmm(&ple.gate, x, d, hidden, GemmAcct::Ffn)?;
let mut gated = vec![0.0f32; seq * d];
for t in 0..seq {
for i in 0..d {
let g = gelu_pytorch_tanh(gate_out[t * d + i]);
let p = ple_tok[t * n_layers * d + li * d + i];
gated[t * d + i] = g * p;
}
}
let proj = self.wmm(&ple.proj, &gated, hidden, d, GemmAcct::Ffn)?;
Ok(Some(self.norm(&proj, &ple.post_norm)?))
}
fn apply_ple(
&self,
x: &mut [f32],
layer: &LayerWeights,
li: usize,
ple_tok: Option<&[f32]>,
hidden: usize,
) -> Result<(), EngineError> {
if let Some(delta) = self.ple_delta(x, layer, li, ple_tok, hidden)? {
for (a, b) in x.iter_mut().zip(delta.iter()) {
*a += *b;
}
}
Ok(())
}
fn apply_layer_scalar(x: &mut [f32], scale: f32) {
if (scale - 1.0).abs() < 1e-8 {
return;
}
for v in x {
*v *= scale;
}
}
fn apply_ffn(
&self,
layer: &LayerWeights,
xn2: &[f32],
hidden: usize,
) -> Result<Vec<f32>, EngineError> {
match &layer.ffn {
FfnWeights::Dense { gate, up, down } => {
if gate.data.len() % hidden != 0 {
return Err(EngineError::ShapeMismatch(
"dense gate len not divisible by hidden".into(),
));
}
let inter = gate.data.len() / hidden;
if inter == 0
|| up.data.len() != inter * hidden
|| down.data.len() != hidden * inter
{
return Err(EngineError::ShapeMismatch(
"dense FFN weight shape mismatch".into(),
));
}
let mut g = self.wmm(gate, xn2, inter, hidden, GemmAcct::Ffn)?;
if layer.activation_sparsity > 0.0 {
g = gaussian_topk(&g, inter, layer.activation_sparsity)?;
}
let u = self.wmm(up, xn2, inter, hidden, GemmAcct::Ffn)?;
let h = if self.use_geglu {
geglu(&g, &u)?
} else {
swiglu(&g, &u)?
};
self.wmm(down, &h, hidden, inter, GemmAcct::Ffn)
}
FfnWeights::MoE {
router,
experts,
top_k,
use_sigmoid,
} => {
let n_exp = experts.len();
let logits = self.wmm(router, xn2, n_exp, hidden, GemmAcct::Ffn)?;
let (ids, weights) = moe_topk_route(&logits, *top_k, *use_sigmoid)?;
let mut acc = vec![0.0f32; hidden];
for (ei, &w) in ids.iter().zip(weights.iter()) {
let ex = &experts[*ei];
if ex.gate.data.len() % hidden != 0 {
return Err(EngineError::ShapeMismatch(
"expert gate len not divisible by hidden".into(),
));
}
let inter = ex.gate.data.len() / hidden;
let g = self.wmm(&ex.gate, xn2, inter, hidden, GemmAcct::Ffn)?;
let u = self.wmm(&ex.up, xn2, inter, hidden, GemmAcct::Ffn)?;
let h = swiglu(&g, &u)?;
let down = self.wmm(&ex.down, &h, hidden, inter, GemmAcct::Ffn)?;
for i in 0..hidden {
acc[i] += w * down[i];
}
}
Ok(acc)
}
}
}
fn fresh_decode_state(&self) -> DecodeState {
let hidden = self.conf.hidden_size;
DecodeState {
k_caches: (0..self.conf.num_layers).map(|_| Vec::new()).collect(),
v_caches: (0..self.conf.num_layers).map(|_| Vec::new()).collect(),
last_kv_src: HashMap::new(),
conv_states: self
.weights
.layers
.iter()
.map(|layer| match &layer.op {
LayerOp::Conv(c) => {
let hist = c.kernel_size.saturating_sub(1);
Some(vec![0.0f32; hidden * hist])
}
LayerOp::Linear(d) => {
let conv_dim = d.n_k_heads * d.head_k * 2 + d.n_v_heads * d.head_v;
let hist = d.conv_k.saturating_sub(1);
Some(vec![0.0f32; conv_dim * hist])
}
LayerOp::Attn(_) => None,
})
.collect(),
delta_states: self
.weights
.layers
.iter()
.map(|layer| match &layer.op {
LayerOp::Linear(d) => Some(vec![0.0f32; d.n_v_heads * d.head_k * d.head_v]),
_ => None,
})
.collect(),
pos: 0,
}
}
#[cfg(test)]
fn forward(&self, tokens: &[u32]) -> Result<Vec<f32>, EngineError> {
let mut state = self.fresh_decode_state();
let mut logits = Vec::new();
for &tok in tokens {
logits = self.forward_step_with(&mut state, tok)?;
}
Ok(logits)
}
fn forward_prompt(&mut self, toks: &[u32]) -> Result<Vec<f32>, EngineError> {
let mut owned = self.decode.take().ok_or_else(|| {
EngineError::InvalidParam("decode state missing; call generate/prefill first".into())
})?;
let logits = self.forward_prompt_with(&mut owned, toks);
self.decode = Some(owned);
logits
}
fn apply_rope_seq(
x: &mut [f32],
seq: usize,
tok_dim: usize,
head_dim: usize,
pos0: usize,
theta: f32,
mode: RopeMode,
) -> Result<(), EngineError> {
if x.len() != seq * tok_dim {
return Err(EngineError::ShapeMismatch(
"rope seq buffer length mismatch".into(),
));
}
for t in 0..seq {
Self::apply_rope(
&mut x[t * tok_dim..(t + 1) * tok_dim],
head_dim,
pos0 + t,
theta,
mode,
)?;
}
Ok(())
}
fn forward_prompt_with(
&self,
state: &mut DecodeState,
toks: &[u32],
) -> Result<Vec<f32>, EngineError> {
if toks.is_empty() {
return Err(EngineError::InvalidParam("empty prompt".into()));
}
let hidden = self.conf.hidden_size;
let n_heads = self.conf.num_attention_heads;
let n_kv = self.conf.num_kv_heads;
let vocab = self.conf.vocab_size;
let seq = toks.len();
if hidden == 0 {
return Err(EngineError::ShapeMismatch("hidden_size is 0".into()));
}
if self.weights.emb.data.len() < vocab.saturating_mul(hidden)
|| !self.weights.emb.data.len().is_multiple_of(hidden)
{
return Err(EngineError::ShapeMismatch(format!(
"embedding length {} not compatible with vocab={vocab} hidden={hidden}",
self.weights.emb.data.len()
)));
}
let pos0 = state.pos;
let mut x = vec![0.0f32; seq * hidden];
for (t, &tok) in toks.iter().enumerate() {
let tid = (tok as usize) % vocab;
x[t * hidden..(t + 1) * hidden]
.copy_from_slice(&self.weights.emb.data[tid * hidden..(tid + 1) * hidden]);
}
if self.embed_scale != 1.0 {
for v in &mut x {
*v *= self.embed_scale;
}
}
let ple_tok = self.compute_ple_inputs(toks, &x)?;
let mut gemma3n_streams = if self.has_gemma3n_graph() {
Some(self.gemma3n_expand_streams(&x, hidden)?)
} else {
None
};
for (li, layer) in self.weights.layers.iter().enumerate() {
let mut gemma3n_preds = None;
let mut gemma3n_laurel = None;
let xn = if let Some(ref streams) = gemma3n_streams {
let preds = self.altup_predict(layer, streams, hidden)?;
let xn = self.norm(&preds[0], &layer.attn_norm)?;
gemma3n_laurel = Some(self.apply_laurel(layer, &xn, hidden)?);
gemma3n_preds = Some(preds);
xn
} else {
self.norm(&x, &layer.attn_norm)?
};
match &layer.op {
LayerOp::Attn(attn) => {
if attn.wq.data.len() % hidden != 0 {
return Err(EngineError::ShapeMismatch(
"attn q proj weight not divisible by hidden_size".into(),
));
}
let q_out = attn.wq.data.len() / hidden;
let q_dim = if attn.q_gate { q_out / 2 } else { q_out };
let head_dim = self.layer_head_dim(attn.kind, q_dim, n_heads)?;
if attn.wo.data.len() != hidden * q_dim {
return Err(EngineError::ShapeMismatch(format!(
"attn output proj weight shape mismatch (wo_len={} hidden={hidden} q_dim={q_dim})",
attn.wo.data.len()
)));
}
let mixed = self.wmm(&attn.wq, &xn, q_out, hidden, GemmAcct::Attn)?;
let (mut q, gate) = if attn.q_gate {
split_interleaved_q_gate(&mixed, seq, n_heads, head_dim)?
} else {
(mixed, Vec::new())
};
if let Some(qn) = &attn.q_norm {
if qn.len() != head_dim {
return Err(EngineError::ShapeMismatch(format!(
"q_norm len {} != head_dim {head_dim}",
qn.len()
)));
}
q = self.norm(&q, qn)?;
}
let (theta, rope) = self.layer_rope_params(attn.kind);
Self::apply_rope_seq(&mut q, seq, q_dim, head_dim, pos0, theta, rope)?;
let (k_src, v_src) = if let (Some(wk), Some(wv)) = (&attn.wk, &attn.wv) {
if wk.data.len() % hidden != 0 || wv.data.len() % hidden != 0 {
return Err(EngineError::ShapeMismatch(
"attn kv proj weight not divisible by hidden_size".into(),
));
}
let k_dim = wk.data.len() / hidden;
let v_dim = wv.data.len() / hidden;
if k_dim != n_kv * head_dim || v_dim != n_kv * head_dim {
return Err(EngineError::ShapeMismatch(format!(
"kv dims {k_dim}/{v_dim} != n_kv*head_dim {}",
n_kv * head_dim
)));
}
let mut k = self.wmm(wk, &xn, k_dim, hidden, GemmAcct::Attn)?;
let mut v = self.wmm(wv, &xn, v_dim, hidden, GemmAcct::Attn)?;
if let Some(kn) = &attn.k_norm {
if kn.len() != head_dim {
return Err(EngineError::ShapeMismatch(format!(
"k_norm len {} != head_dim {head_dim}",
kn.len()
)));
}
k = self.norm(&k, kn)?;
}
Self::apply_rope_seq(&mut k, seq, k_dim, head_dim, pos0, theta, rope)?;
v = self.apply_v_norm(v, attn.v_norm.as_deref(), head_dim)?;
state.k_caches[li] = k;
state.v_caches[li] = v;
state.last_kv_src.insert(attn.kind, li);
(li, li)
} else {
let src = state.last_kv_src.get(&attn.kind).copied().ok_or_else(|| {
EngineError::Format(format!(
"KV-consumer layer {li} has no producer of kind {:?}",
attn.kind
))
})?;
(src, src)
};
let attn_out = attention_causal_with_scale(
&q,
&state.k_caches[k_src],
&state.v_caches[v_src],
n_heads,
n_kv,
head_dim,
self.attn_scale(head_dim),
self.attn_window(attn.kind),
)?;
let mut attn_out = attn_out;
if attn.q_gate {
apply_sigmoid_gate(&mut attn_out, &gate)?;
}
let ao = self.wmm(&attn.wo, &attn_out, hidden, q_dim, GemmAcct::Attn)?;
if let (Some(ref mut streams), Some(preds), Some(laurel)) =
(gemma3n_streams.as_mut(), gemma3n_preds, gemma3n_laurel)
{
self.gemma3n_after_attn(
streams,
&preds,
&laurel,
&ao,
layer,
li,
ple_tok.as_deref(),
hidden,
)?;
continue;
}
self.add_normed_residual(&mut x, &ao, layer.post_attn_norm.as_deref())?;
}
LayerOp::Conv(_) | LayerOp::Linear(_) => {
return Err(EngineError::Unsupported(
"batched prefill is only implemented for attention+dense FFN layers".into(),
));
}
}
let xn2 = self.norm(&x, &layer.ffn_norm)?;
let down = self.apply_ffn(layer, &xn2, hidden)?;
self.add_normed_residual(&mut x, &down, layer.post_ffn_norm.as_deref())?;
self.apply_ple(&mut x, layer, li, ple_tok.as_deref(), hidden)?;
Self::apply_layer_scalar(&mut x, layer.layer_scalar);
}
if let Some(streams) = gemma3n_streams {
x = self.gemma3n_unembed_streams(&streams, hidden)?;
}
state.pos = pos0 + seq;
let last = &x[(seq - 1) * hidden..seq * hidden];
let xn = self.norm(last, &self.weights.output_norm)?;
if !self.weights.output.data.len().is_multiple_of(hidden) {
return Err(EngineError::ShapeMismatch(format!(
"lm_head len {} not divisible by hidden {hidden}",
self.weights.output.data.len()
)));
}
let out_rows = self.weights.output.data.len() / hidden;
let logits = self.wmm(
&self.weights.output,
&xn,
out_rows,
hidden,
GemmAcct::LmHead,
)?;
Ok(self.softcap_logits(logits))
}
fn forward_step(&mut self, tok: u32) -> Result<Vec<f32>, EngineError> {
let mut owned = self.decode.take().ok_or_else(|| {
EngineError::InvalidParam("decode state missing; call generate/prefill first".into())
})?;
let logits = self.forward_step_with(&mut owned, tok);
self.decode = Some(owned);
logits
}
fn forward_step_with(
&self,
state: &mut DecodeState,
tok: u32,
) -> Result<Vec<f32>, EngineError> {
let hidden = self.conf.hidden_size;
let n_heads = self.conf.num_attention_heads;
let n_kv = self.conf.num_kv_heads;
let vocab = self.conf.vocab_size;
if hidden == 0 {
return Err(EngineError::ShapeMismatch("hidden_size is 0".into()));
}
if self.weights.emb.data.len() < vocab.saturating_mul(hidden)
|| !self.weights.emb.data.len().is_multiple_of(hidden)
{
return Err(EngineError::ShapeMismatch(format!(
"embedding length {} not compatible with vocab={vocab} hidden={hidden}",
self.weights.emb.data.len()
)));
}
let pos = state.pos;
let tid = (tok as usize) % vocab;
let mut x = vec![0.0f32; hidden];
x.copy_from_slice(&self.weights.emb.data[tid * hidden..(tid + 1) * hidden]);
if self.embed_scale != 1.0 {
for v in &mut x {
*v *= self.embed_scale;
}
}
let ple_tok = self.compute_ple_inputs(&[tok], &x)?;
let mut gemma3n_streams = if self.has_gemma3n_graph() {
Some(self.gemma3n_expand_streams(&x, hidden)?)
} else {
None
};
for (li, layer) in self.weights.layers.iter().enumerate() {
let mut gemma3n_preds = None;
let mut gemma3n_laurel = None;
let xn = if let Some(ref streams) = gemma3n_streams {
let preds = self.altup_predict(layer, streams, hidden)?;
let xn = self.norm(&preds[0], &layer.attn_norm)?;
gemma3n_laurel = Some(self.apply_laurel(layer, &xn, hidden)?);
gemma3n_preds = Some(preds);
xn
} else {
self.norm(&x, &layer.attn_norm)?
};
match &layer.op {
LayerOp::Attn(attn) => {
if attn.wq.data.len() % hidden != 0 {
return Err(EngineError::ShapeMismatch(
"attn q proj weight not divisible by hidden_size".into(),
));
}
let q_out = attn.wq.data.len() / hidden;
let q_dim = if attn.q_gate { q_out / 2 } else { q_out };
let head_dim = self.layer_head_dim(attn.kind, q_dim, n_heads)?;
if attn.wo.data.len() != hidden * q_dim {
return Err(EngineError::ShapeMismatch(format!(
"attn output proj weight shape mismatch (wo_len={} hidden={hidden} q_dim={q_dim})",
attn.wo.data.len()
)));
}
let mixed = self.wmm(&attn.wq, &xn, q_out, hidden, GemmAcct::Attn)?;
let (mut q, gate) = if attn.q_gate {
split_interleaved_q_gate(&mixed, 1, n_heads, head_dim)?
} else {
(mixed, Vec::new())
};
if let Some(qn) = &attn.q_norm {
if qn.len() != head_dim {
return Err(EngineError::ShapeMismatch(format!(
"q_norm len {} != head_dim {head_dim}",
qn.len()
)));
}
q = self.norm(&q, qn)?;
}
let (theta, rope) = self.layer_rope_params(attn.kind);
Self::apply_rope(&mut q, head_dim, pos, theta, rope)?;
let (k_src, v_src) = if let (Some(wk), Some(wv)) = (&attn.wk, &attn.wv) {
if wk.data.len() % hidden != 0 || wv.data.len() % hidden != 0 {
return Err(EngineError::ShapeMismatch(
"attn kv proj weight not divisible by hidden_size".into(),
));
}
let k_dim = wk.data.len() / hidden;
let v_dim = wv.data.len() / hidden;
if k_dim != n_kv * head_dim || v_dim != n_kv * head_dim {
return Err(EngineError::ShapeMismatch(format!(
"kv dims {k_dim}/{v_dim} != n_kv*head_dim {}",
n_kv * head_dim
)));
}
let mut k = self.wmm(wk, &xn, k_dim, hidden, GemmAcct::Attn)?;
let mut v = self.wmm(wv, &xn, v_dim, hidden, GemmAcct::Attn)?;
if let Some(kn) = &attn.k_norm {
if kn.len() != head_dim {
return Err(EngineError::ShapeMismatch(format!(
"k_norm len {} != head_dim {head_dim}",
kn.len()
)));
}
k = self.norm(&k, kn)?;
}
Self::apply_rope(&mut k, head_dim, pos, theta, rope)?;
v = self.apply_v_norm(v, attn.v_norm.as_deref(), head_dim)?;
state.k_caches[li].extend_from_slice(&k);
state.v_caches[li].extend_from_slice(&v);
state.last_kv_src.insert(attn.kind, li);
(li, li)
} else {
let src = state.last_kv_src.get(&attn.kind).copied().ok_or_else(|| {
EngineError::Format(format!(
"KV-consumer layer {li} has no producer of kind {:?}",
attn.kind
))
})?;
(src, src)
};
let kv_dim = n_kv * head_dim;
let (k_view, v_view) = kv_sliding_view(
&state.k_caches[k_src],
&state.v_caches[v_src],
kv_dim,
self.attn_window(attn.kind),
)?;
let attn_out = attention_with_scale(
&q,
k_view,
v_view,
n_heads,
n_kv,
head_dim,
self.attn_scale(head_dim),
)?;
let mut attn_out = attn_out;
if attn.q_gate {
apply_sigmoid_gate(&mut attn_out, &gate)?;
}
let ao = self.wmm(&attn.wo, &attn_out, hidden, q_dim, GemmAcct::Attn)?;
if let (Some(ref mut streams), Some(preds), Some(laurel)) =
(gemma3n_streams.as_mut(), gemma3n_preds, gemma3n_laurel)
{
self.gemma3n_after_attn(
streams,
&preds,
&laurel,
&ao,
layer,
li,
ple_tok.as_deref(),
hidden,
)?;
continue;
}
self.add_normed_residual(&mut x, &ao, layer.post_attn_norm.as_deref())?;
}
LayerOp::Conv(conv) => {
let bcx = self.wmm(&conv.in_proj, &xn, 3 * hidden, hidden, GemmAcct::Attn)?;
let mut bx = vec![0.0f32; hidden];
let mut c_gate = vec![0.0f32; hidden];
for i in 0..hidden {
let b = bcx[i];
let c = bcx[hidden + i];
let xx = bcx[2 * hidden + i];
bx[i] = b * xx;
c_gate[i] = c;
}
let cstate = state.conv_states[li]
.as_mut()
.ok_or_else(|| EngineError::ShapeMismatch("missing conv state".into()))?;
let conv_y =
short_conv_step(&bx, &conv.kernel, cstate, hidden, conv.kernel_size)?;
let mut y = vec![0.0f32; hidden];
for i in 0..hidden {
y[i] = c_gate[i] * conv_y[i];
}
let ao = self.wmm(&conv.out_proj, &y, hidden, hidden, GemmAcct::Attn)?;
self.add_normed_residual(&mut x, &ao, layer.post_attn_norm.as_deref())?;
}
LayerOp::Linear(dn) => {
let key_dim = dn.n_k_heads * dn.head_k;
let value_dim = dn.n_v_heads * dn.head_v;
let qkvz_out = 2 * key_dim + 2 * value_dim;
let mixed = self.wmm(&dn.qkvz, &xn, qkvz_out, hidden, GemmAcct::Attn)?;
let mut q = mixed[0..key_dim].to_vec();
let mut k = mixed[key_dim..2 * key_dim].to_vec();
let mut v = mixed[2 * key_dim..2 * key_dim + value_dim].to_vec();
let z = mixed[2 * key_dim + value_dim..].to_vec();
let mut qkv = Vec::with_capacity(key_dim * 2 + value_dim);
qkv.extend_from_slice(&q);
qkv.extend_from_slice(&k);
qkv.extend_from_slice(&v);
let conv_dim = qkv.len();
let cstate = state.conv_states[li].as_mut().ok_or_else(|| {
EngineError::ShapeMismatch("missing delta conv state".into())
})?;
let mut mixed_c = short_conv_step(&qkv, &dn.conv, cstate, conv_dim, dn.conv_k)?;
silu_vec(&mut mixed_c);
q.copy_from_slice(&mixed_c[0..key_dim]);
k.copy_from_slice(&mixed_c[key_dim..2 * key_dim]);
v.copy_from_slice(&mixed_c[2 * key_dim..]);
let ba = self.wmm(&dn.ba, &xn, 2 * dn.n_v_heads, hidden, GemmAcct::Attn)?;
let mut beta = vec![0.0f32; dn.n_v_heads];
let mut g = vec![0.0f32; dn.n_v_heads];
for h in 0..dn.n_v_heads {
beta[h] = 1.0 / (1.0 + (-ba[h]).exp());
let alpha =
-dn.a_log[h].exp() * softplus(ba[dn.n_v_heads + h] + dn.dt_bias[h]);
g[h] = alpha.exp();
}
if dn.n_v_heads != dn.n_k_heads {
return Err(EngineError::Unsupported(
"DeltaNet GQA (n_v != n_k) not implemented".into(),
));
}
let s = state.delta_states[li].as_mut().ok_or_else(|| {
EngineError::ShapeMismatch("missing delta recurrent state".into())
})?;
let mut core = gated_delta_step(GatedDeltaStep {
q: &q,
k: &k,
v: &v,
g: &g,
beta: &beta,
state: s,
n_heads: dn.n_v_heads,
dk: dn.head_k,
dv: dn.head_v,
})?;
core = rms_norm(&core, &dn.out_norm, 1e-6)?;
let mut z_act = z;
silu_vec(&mut z_act);
for i in 0..core.len() {
core[i] *= z_act[i];
}
let ao = self.wmm(&dn.out_proj, &core, hidden, value_dim, GemmAcct::Attn)?;
self.add_normed_residual(&mut x, &ao, layer.post_attn_norm.as_deref())?;
}
}
let xn2 = self.norm(&x, &layer.ffn_norm)?;
let down = self.apply_ffn(layer, &xn2, hidden)?;
self.add_normed_residual(&mut x, &down, layer.post_ffn_norm.as_deref())?;
self.apply_ple(&mut x, layer, li, ple_tok.as_deref(), hidden)?;
Self::apply_layer_scalar(&mut x, layer.layer_scalar);
}
if let Some(streams) = gemma3n_streams {
x = self.gemma3n_unembed_streams(&streams, hidden)?;
}
state.pos += 1;
let xn = self.norm(&x, &self.weights.output_norm)?;
if !self.weights.output.data.len().is_multiple_of(hidden) {
return Err(EngineError::ShapeMismatch(format!(
"lm_head len {} not divisible by hidden {hidden}",
self.weights.output.data.len()
)));
}
let out_rows = self.weights.output.data.len() / hidden;
let logits = self.wmm(
&self.weights.output,
&xn,
out_rows,
hidden,
GemmAcct::LmHead,
)?;
Ok(self.softcap_logits(logits))
}
fn softcap_logits(&self, mut logits: Vec<f32>) -> Vec<f32> {
if let Some(cap) = self.final_logit_softcap.filter(|c| *c > 0.0) {
for x in &mut logits {
*x = (*x / cap).tanh() * cap;
}
}
logits
}
}
fn gaussian_topk(gate: &[f32], inter: usize, sparsity: f32) -> Result<Vec<f32>, EngineError> {
if inter == 0 || !gate.len().is_multiple_of(inter) {
return Err(EngineError::ShapeMismatch(
"gaussian_topk: gate len not divisible by intermediate_size".into(),
));
}
let z = if (sparsity - 0.95).abs() < 0.02 {
1.644_853_8
} else {
1.644_853_8 * (sparsity / 0.95).clamp(0.0, 4.0)
};
let seq = gate.len() / inter;
let mut out = vec![0.0f32; gate.len()];
let n = inter as f32;
for t in 0..seq {
let row = &gate[t * inter..(t + 1) * inter];
let mean = row.iter().sum::<f32>() / n;
let mut var = 0.0f32;
for &v in row {
let d = v - mean;
var += d * d;
}
var /= n;
let cutoff = mean + var.sqrt() * z;
for i in 0..inter {
out[t * inter + i] = (row[i] - cutoff).max(0.0);
}
}
Ok(out)
}
fn argmax(v: &[f32]) -> u32 {
let mut best = 0usize;
let mut best_v = f32::NEG_INFINITY;
for (i, &x) in v.iter().enumerate() {
if x > best_v {
best_v = x;
best = i;
}
}
best as u32
}
pub fn confidence_from_logits(logits: &[f32]) -> f32 {
if logits.is_empty() {
return 0.0;
}
let m = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0f32;
let mut maxp = 0.0f32;
for &x in logits {
let e = (x - m).exp();
sum += e;
if e > maxp {
maxp = e;
}
}
if sum > 0.0 {
maxp / sum
} else {
0.0
}
}
#[allow(dead_code)]
pub fn cache_shapes_ok(cache: &HashMap<usize, Vec<f32>>, kv_dim: usize) -> bool {
cache.values().all(|v| v.len().is_multiple_of(kv_dim))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::family::{arch_class_representatives, graph_hook, lookup_family, require_stage_b};
use crate::fixture::write_tiny_q4_bundle;
use aria_kernel::{resolve_compute, ComputePref};
use serde_json::{json, Value};
#[test]
fn gemma4_fills_hub_bundle_missing_geometry_fields() {
let dir = tempfile::tempdir().unwrap();
write_tiny_q4_bundle(dir.path()).unwrap();
let cfg_path = dir.path().join("config.json");
let raw = std::fs::read_to_string(&cfg_path).unwrap();
let mut cfg: Value = serde_json::from_str(&raw).unwrap();
let model = cfg["model"].as_object_mut().unwrap();
for key in [
"layer_types",
"sliding_window",
"partial_rotary_factor",
"global_head_dim",
"head_dim",
] {
model.remove(key);
}
std::fs::write(&cfg_path, cfg.to_string()).unwrap();
let s = SessionBuilder::new()
.model(dir.path())
.family("gemma/gemma-4-e2b-it")
.build()
.unwrap();
assert_eq!(s.config().sliding_window, Some(512));
assert_eq!(s.config().partial_rotary_factor, Some(0.25));
assert!(s.config().head_dim.unwrap_or(0) > 0);
assert!(s.config().global_head_dim.unwrap_or(0) > 0);
assert_eq!(
s.config().layer_types.as_ref().map(|t| t.len()),
Some(s.config().num_layers)
);
}
#[test]
fn gemma3_text_not_gemma3n() {
assert!(is_gemma3_text("gemma/gemma-3-270m-it"));
assert!(is_gemma3_text("gemma/gemma-3-1b-it"));
assert!(!is_gemma3_text("gemma/gemma-3n-e2b-it"));
assert!(!is_gemma3_text("gemma/gemma-4-e2b-it"));
let t = default_gemma3_layer_types(18);
assert_eq!(t[5], "full_attention");
assert_eq!(t[11], "full_attention");
assert_eq!(t[17], "full_attention");
assert_eq!(t.iter().filter(|s| *s == "sliding_attention").count(), 15);
}
#[test]
fn gemma3_fills_hub_bundle_and_dual_rope() {
let dir = tempfile::tempdir().unwrap();
write_tiny_q4_bundle(dir.path()).unwrap();
let cfg_path = dir.path().join("config.json");
let raw = std::fs::read_to_string(&cfg_path).unwrap();
let mut cfg: Value = serde_json::from_str(&raw).unwrap();
let model = cfg["model"].as_object_mut().unwrap();
for key in ["layer_types", "sliding_window", "hidden_act"] {
model.remove(key);
}
std::fs::write(&cfg_path, cfg.to_string()).unwrap();
let mut s = SessionBuilder::new()
.model(dir.path())
.family("gemma/gemma-3-270m-it")
.build()
.unwrap();
assert_eq!(s.config().sliding_window, Some(512));
assert_eq!(s.config().hidden_act.as_deref(), Some("gelu_pytorch_tanh"));
let types = s.config().layer_types.as_ref().expect("layer_types");
assert_eq!(types.len(), s.config().num_layers);
assert!(types.iter().all(|t| t == "sliding_attention"));
assert_eq!(
s.layer_rope_params(AttnKind::Sliding),
(10_000.0, RopeMode::Full)
);
assert_eq!(
s.layer_rope_params(AttnKind::Full),
(1_000_000.0, RopeMode::Full)
);
let gen = s
.generate(
&[1, 2],
&GenerateOpts {
max_tokens: 2,
temperature: 0.0,
},
)
.unwrap();
assert_eq!(gen.tokens.len(), 2);
}
#[test]
fn gemma3n_not_gemma3_text_and_fills_4plus1_dual_rope() {
assert!(is_gemma3n("gemma/gemma-3n-e2b-it"));
assert!(is_gemma3n("gemma/gemma-3n-e4b-it"));
assert!(!is_gemma3n("gemma/gemma-3-270m-it"));
let t = default_gemma4_layer_types(30);
assert_eq!(t[4], "full_attention");
assert_eq!(t[9], "full_attention");
assert_eq!(t[29], "full_attention");
assert_eq!(t.iter().filter(|s| *s == "sliding_attention").count(), 24);
assert_eq!(gemma3n_default_kv_shared(30), 10);
assert_eq!(gemma3n_default_kv_shared(35), 15);
assert_eq!(gemma3n_default_kv_shared(2), 0);
assert!(
(altup_router_input_scale(2048) - 1.0 / 2048.0).abs() < 1e-12,
"HF router_input_scale is 1/hidden, not 1/sqrt(hidden)"
);
let dir = tempfile::tempdir().unwrap();
write_tiny_q4_bundle(dir.path()).unwrap();
let cfg_path = dir.path().join("config.json");
let raw = std::fs::read_to_string(&cfg_path).unwrap();
let mut cfg: Value = serde_json::from_str(&raw).unwrap();
let model = cfg["model"].as_object_mut().unwrap();
for key in ["layer_types", "sliding_window", "hidden_act"] {
model.remove(key);
}
std::fs::write(&cfg_path, cfg.to_string()).unwrap();
let mut s = SessionBuilder::new()
.model(dir.path())
.family("gemma/gemma-3n-e2b-it")
.build()
.unwrap();
assert!(!s.use_gemma_norm, "Gemma-3n RMSNorm is *w, not *(1+w)");
assert!((s.attn_scale(256) - 1.0).abs() < 1e-6);
assert_eq!(s.final_logit_softcap, Some(30.0));
assert_eq!(s.config().sliding_window, Some(512));
assert_eq!(
s.layer_rope_params(AttnKind::Sliding),
(10_000.0, RopeMode::Full)
);
assert_eq!(
s.layer_rope_params(AttnKind::Full),
(1_000_000.0, RopeMode::Full)
);
let gen = s
.generate(
&[1, 2],
&GenerateOpts {
max_tokens: 2,
temperature: 0.0,
},
)
.unwrap();
assert_eq!(gen.tokens.len(), 2);
}
#[test]
fn generate_tokens() {
let dir = tempfile::tempdir().unwrap();
write_tiny_q4_bundle(dir.path()).unwrap();
let mut s = SessionBuilder::new()
.model(dir.path())
.family("gemma/gemma-4-e2b-it")
.build()
.unwrap();
assert_eq!(s.config().sliding_window, Some(512));
assert_eq!(s.config().partial_rotary_factor, Some(0.25));
assert_eq!(s.config().head_dim, Some(16));
assert_eq!(s.config().global_head_dim, Some(16));
assert_eq!(
s.config().layer_types,
Some(vec!["full_attention".into(), "full_attention".into()])
);
assert_eq!(
s.layer_rope_params(AttnKind::Full),
(1_000_000.0, RopeMode::Proportional(0.25))
);
assert_eq!(
s.layer_rope_params(AttnKind::Sliding),
(10_000.0, RopeMode::Full)
);
assert_eq!(s.attn_window(AttnKind::Sliding), Some(512));
assert_eq!(s.attn_window(AttnKind::Full), None);
let prompt = s.encode_text("hi");
let gen = s
.generate(
&prompt,
&GenerateOpts {
max_tokens: 4,
temperature: 0.0,
},
)
.unwrap();
assert!(!gen.tokens.is_empty());
assert!(!gen.text.is_empty());
}
#[test]
fn split_interleaved_q_gate_matches_hf_chunk() {
let mixed = vec![1.0, 2.0, 10.0, 20.0, 3.0, 4.0, 30.0, 40.0];
let (q, g) = split_interleaved_q_gate(&mixed, 1, 2, 2).unwrap();
assert_eq!(q, vec![1.0, 2.0, 3.0, 4.0]);
assert_eq!(g, vec![10.0, 20.0, 30.0, 40.0]);
}
#[test]
fn qwen35_attn_output_gate_generate() {
let dir = tempfile::tempdir().unwrap();
let hidden = 8usize;
let inter = 16usize;
let vocab = 16usize;
let n_heads = 2usize;
let n_kv = 1usize;
let head_dim = 4usize;
let q_dim = n_heads * head_dim;
let k_dim = n_kv * head_dim;
let mut tensors = serde_json::Map::new();
let mut bin = Vec::new();
let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
let offset = bin.len();
for &v in data {
bin.extend_from_slice(&v.to_le_bytes());
}
let nbytes = data.len() * 4;
let mut meta = serde_json::Map::new();
meta.insert("kind".into(), json!("raw"));
meta.insert("dtype".into(), json!("f32"));
meta.insert("shape".into(), json!(shape));
meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
tensors.insert(name.to_string(), Value::Object(meta));
};
let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
let n1 = vec![1.0f32; hidden];
add_raw("model.layers.0.input_layernorm.weight", vec![hidden], &n1);
add_raw(
"model.layers.0.post_attention_layernorm.weight",
vec![hidden],
&n1,
);
let wq = vec![0.02f32; 2 * q_dim * hidden];
let wk = vec![0.02f32; k_dim * hidden];
let wv = vec![0.02f32; k_dim * hidden];
let wo = vec![0.02f32; hidden * q_dim];
add_raw(
"model.layers.0.self_attn.q_proj.weight",
vec![2 * q_dim, hidden],
&wq,
);
add_raw(
"model.layers.0.self_attn.k_proj.weight",
vec![k_dim, hidden],
&wk,
);
add_raw(
"model.layers.0.self_attn.v_proj.weight",
vec![k_dim, hidden],
&wv,
);
add_raw(
"model.layers.0.self_attn.o_proj.weight",
vec![hidden, q_dim],
&wo,
);
let qn = vec![1.0f32; head_dim];
add_raw(
"model.layers.0.self_attn.q_norm.weight",
vec![head_dim],
&qn,
);
add_raw(
"model.layers.0.self_attn.k_norm.weight",
vec![head_dim],
&qn,
);
let g = vec![0.02f32; inter * hidden];
let d = vec![0.02f32; hidden * inter];
add_raw(
"model.layers.0.mlp.gate_proj.weight",
vec![inter, hidden],
&g,
);
add_raw("model.layers.0.mlp.up_proj.weight", vec![inter, hidden], &g);
add_raw(
"model.layers.0.mlp.down_proj.weight",
vec![hidden, inter],
&d,
);
add_raw("model.norm.weight", vec![hidden], &n1);
add_raw("lm_head.weight", vec![vocab, hidden], &emb);
let cfg = json!({
"format": "aria-quant-bundle",
"format_version": 2,
"quantization": "test",
"hadamard_seed": 0,
"model": {
"hidden_size": hidden,
"num_layers": 1,
"num_attention_heads": n_heads,
"num_kv_heads": n_kv,
"head_dim": head_dim,
"intermediate_size": inter,
"vocab_size": vocab,
"context_length": 32,
"rope_theta": 10000.0,
"layer_types": ["full_attention"]
},
"tensors": tensors
});
std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
let mut s = SessionBuilder::new()
.model(dir.path())
.family("qwen/qwen3-0.6b")
.build()
.unwrap();
let gen = s
.generate(
&[1, 2],
&GenerateOpts {
max_tokens: 2,
temperature: 0.0,
},
)
.unwrap();
assert_eq!(gen.tokens.len(), 2);
}
#[test]
fn materialize_accepts_hf_tensor_names() {
let dir = tempfile::tempdir().unwrap();
let hidden = 8usize;
let layers = 1usize;
let inter = 16usize;
let vocab = 16usize;
let n_heads = 2usize;
let n_kv = 1usize;
let head_dim = 4usize; let q_dim = n_heads * head_dim;
let k_dim = n_kv * head_dim;
let mut tensors = serde_json::Map::new();
let mut bin = Vec::new();
let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
let offset = bin.len();
for &v in data {
bin.extend_from_slice(&v.to_le_bytes());
}
let nbytes = data.len() * 4;
let mut meta = serde_json::Map::new();
meta.insert("kind".into(), json!("raw"));
meta.insert("dtype".into(), json!("f32"));
meta.insert("shape".into(), json!(shape));
meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
tensors.insert(name.to_string(), Value::Object(meta));
};
let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
let n1 = vec![1.0f32; hidden];
add_raw("model.layers.0.input_layernorm.weight", vec![hidden], &n1);
add_raw(
"model.layers.0.post_attention_layernorm.weight",
vec![hidden],
&n1,
);
let wq = vec![0.01f32; q_dim * hidden];
let wk = vec![0.01f32; k_dim * hidden];
let wv = vec![0.01f32; k_dim * hidden];
let wo = vec![0.01f32; hidden * q_dim];
add_raw(
"model.layers.0.self_attn.q_proj.weight",
vec![q_dim, hidden],
&wq,
);
add_raw(
"model.layers.0.self_attn.k_proj.weight",
vec![k_dim, hidden],
&wk,
);
add_raw(
"model.layers.0.self_attn.v_proj.weight",
vec![k_dim, hidden],
&wv,
);
add_raw(
"model.layers.0.self_attn.o_proj.weight",
vec![hidden, q_dim],
&wo,
);
let g = vec![0.01f32; inter * hidden];
let d = vec![0.01f32; hidden * inter];
add_raw(
"model.layers.0.mlp.gate_proj.weight",
vec![inter, hidden],
&g,
);
add_raw("model.layers.0.mlp.up_proj.weight", vec![inter, hidden], &g);
add_raw(
"model.layers.0.mlp.down_proj.weight",
vec![hidden, inter],
&d,
);
add_raw("model.norm.weight", vec![hidden], &n1);
add_raw("lm_head.weight", vec![vocab, hidden], &emb);
let cfg = json!({
"format": "aria-quant-bundle",
"format_version": 2,
"quantization": "test",
"group_size_default": 32,
"hadamard_seed": 0,
"model": {
"hidden_size": hidden,
"num_layers": layers,
"num_attention_heads": n_heads,
"num_kv_heads": n_kv,
"intermediate_size": inter,
"vocab_size": vocab,
"context_length": 32,
"rope_theta": 10000.0
},
"tensors": tensors
});
std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
let mut s = SessionBuilder::new()
.model(dir.path())
.family("qwen/qwen3-0.6b")
.build()
.unwrap();
let gen = s
.generate(
&[1, 2],
&GenerateOpts {
max_tokens: 2,
temperature: 0.0,
},
)
.unwrap();
assert_eq!(gen.tokens.len(), 2);
}
#[test]
fn materialize_accepts_language_model_prefix_and_pre_ffn_norm() {
let dir = tempfile::tempdir().unwrap();
let hidden = 8usize;
let layers = 1usize;
let inter = 16usize;
let vocab = 16usize;
let n_heads = 2usize;
let n_kv = 1usize;
let head_dim = 4usize;
let q_dim = n_heads * head_dim;
let k_dim = n_kv * head_dim;
let mut tensors = serde_json::Map::new();
let mut bin = Vec::new();
let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
let offset = bin.len();
for &v in data {
bin.extend_from_slice(&v.to_le_bytes());
}
let nbytes = data.len() * 4;
let mut meta = serde_json::Map::new();
meta.insert("kind".into(), json!("raw"));
meta.insert("dtype".into(), json!("f32"));
meta.insert("shape".into(), json!(shape));
meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
tensors.insert(name.to_string(), Value::Object(meta));
};
let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
let p = "model.language_model";
add_raw(
&format!("{p}.embed_tokens.weight"),
vec![vocab, hidden],
&emb,
);
let n1 = vec![1.0f32; hidden];
add_raw(
&format!("{p}.layers.0.input_layernorm.weight"),
vec![hidden],
&n1,
);
add_raw(
&format!("{p}.layers.0.pre_feedforward_layernorm.weight"),
vec![hidden],
&n1,
);
let wq = vec![0.01f32; q_dim * hidden];
let wk = vec![0.01f32; k_dim * hidden];
let wv = vec![0.01f32; k_dim * hidden];
let wo = vec![0.01f32; hidden * q_dim];
add_raw(
&format!("{p}.layers.0.self_attn.q_proj.weight"),
vec![q_dim, hidden],
&wq,
);
add_raw(
&format!("{p}.layers.0.self_attn.k_proj.weight"),
vec![k_dim, hidden],
&wk,
);
add_raw(
&format!("{p}.layers.0.self_attn.v_proj.weight"),
vec![k_dim, hidden],
&wv,
);
add_raw(
&format!("{p}.layers.0.self_attn.o_proj.weight"),
vec![hidden, q_dim],
&wo,
);
let g = vec![0.01f32; inter * hidden];
let d = vec![0.01f32; hidden * inter];
add_raw(
&format!("{p}.layers.0.mlp.gate_proj.weight"),
vec![inter, hidden],
&g,
);
add_raw(
&format!("{p}.layers.0.mlp.up_proj.weight"),
vec![inter, hidden],
&g,
);
add_raw(
&format!("{p}.layers.0.mlp.down_proj.weight"),
vec![hidden, inter],
&d,
);
add_raw(&format!("{p}.norm.weight"), vec![hidden], &n1);
add_raw("lm_head.weight", vec![vocab, hidden], &emb);
let cfg = json!({
"format": "aria-quant-bundle",
"format_version": 2,
"quantization": "test",
"group_size_default": 32,
"hadamard_seed": 0,
"model": {
"hidden_size": hidden,
"num_layers": layers,
"num_attention_heads": n_heads,
"num_kv_heads": n_kv,
"intermediate_size": inter,
"vocab_size": vocab,
"context_length": 32,
"rope_theta": 10000.0,
"head_dim": head_dim,
"global_head_dim": head_dim,
"sliding_window": 512,
"partial_rotary_factor": 0.25,
"layer_types": ["full_attention"]
},
"tensors": tensors
});
std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
let mut s = SessionBuilder::new()
.model(dir.path())
.family("gemma/gemma-4-e2b-it")
.build()
.unwrap();
let gen = s
.generate(
&[1, 2],
&GenerateOpts {
max_tokens: 2,
temperature: 0.0,
},
)
.unwrap();
assert_eq!(gen.tokens.len(), 2);
}
#[test]
fn gemma4_style_double_wide_mlp_and_shared_kv() {
let dir = tempfile::tempdir().unwrap();
let hidden = 8usize;
let layers = 2usize;
let inter = 16usize;
let inter_wide = 32usize;
let vocab = 16usize;
let n_heads = 2usize;
let n_kv = 1usize;
let head_dim = 4usize;
let q_dim = n_heads * head_dim;
let k_dim = n_kv * head_dim;
let p = "model.language_model";
let mut tensors = serde_json::Map::new();
let mut bin = Vec::new();
let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
let offset = bin.len();
for &v in data {
bin.extend_from_slice(&v.to_le_bytes());
}
let nbytes = data.len() * 4;
let mut meta = serde_json::Map::new();
meta.insert("kind".into(), json!("raw"));
meta.insert("dtype".into(), json!("f32"));
meta.insert("shape".into(), json!(shape));
meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
tensors.insert(name.to_string(), Value::Object(meta));
};
let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
add_raw(
&format!("{p}.embed_tokens.weight"),
vec![vocab, hidden],
&emb,
);
let n1 = vec![1.0f32; hidden];
let wq = vec![0.01f32; q_dim * hidden];
let wk = vec![0.01f32; k_dim * hidden];
let wv = vec![0.01f32; k_dim * hidden];
let wo = vec![0.01f32; hidden * q_dim];
for li in 0..layers {
let layer_inter = if li == 0 { inter } else { inter_wide };
add_raw(
&format!("{p}.layers.{li}.input_layernorm.weight"),
vec![hidden],
&n1,
);
add_raw(
&format!("{p}.layers.{li}.pre_feedforward_layernorm.weight"),
vec![hidden],
&n1,
);
add_raw(
&format!("{p}.layers.{li}.self_attn.q_proj.weight"),
vec![q_dim, hidden],
&wq,
);
if li == 0 {
add_raw(
&format!("{p}.layers.{li}.self_attn.k_proj.weight"),
vec![k_dim, hidden],
&wk,
);
add_raw(
&format!("{p}.layers.{li}.self_attn.v_proj.weight"),
vec![k_dim, hidden],
&wv,
);
}
add_raw(
&format!("{p}.layers.{li}.self_attn.o_proj.weight"),
vec![hidden, q_dim],
&wo,
);
let g = vec![0.01f32; layer_inter * hidden];
let d = vec![0.01f32; hidden * layer_inter];
add_raw(
&format!("{p}.layers.{li}.mlp.gate_proj.weight"),
vec![layer_inter, hidden],
&g,
);
add_raw(
&format!("{p}.layers.{li}.mlp.up_proj.weight"),
vec![layer_inter, hidden],
&g,
);
add_raw(
&format!("{p}.layers.{li}.mlp.down_proj.weight"),
vec![hidden, layer_inter],
&d,
);
}
add_raw(&format!("{p}.norm.weight"), vec![hidden], &n1);
add_raw("lm_head.weight", vec![vocab, hidden], &emb);
let cfg = json!({
"format": "aria-quant-bundle",
"format_version": 2,
"quantization": "test",
"group_size_default": 32,
"hadamard_seed": 0,
"model": {
"hidden_size": hidden,
"num_layers": layers,
"num_attention_heads": n_heads,
"num_kv_heads": n_kv,
"intermediate_size": inter,
"vocab_size": vocab,
"context_length": 32,
"rope_theta": 10000.0,
"num_kv_shared_layers": 1,
"head_dim": head_dim,
"global_head_dim": head_dim,
"sliding_window": 512,
"partial_rotary_factor": 0.25,
"layer_types": ["full_attention", "full_attention"]
},
"tensors": tensors
});
std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
let mut s = SessionBuilder::new()
.model(dir.path())
.family("gemma/gemma-4-e2b-it")
.build()
.unwrap();
let gen = s
.generate(
&[1, 2],
&GenerateOpts {
max_tokens: 2,
temperature: 0.0,
},
)
.unwrap();
assert_eq!(gen.tokens.len(), 2);
}
#[test]
fn stage_b_arch_classes_generate() {
for (path, arch) in arch_class_representatives() {
if matches!(arch, ArchClass::VL | ArchClass::VLA | ArchClass::TextMoE) {
continue; }
if path.contains("qwen3.5") || path.contains("bonsai") {
let dir = tempfile::tempdir().unwrap();
write_tiny_q4_bundle(dir.path()).unwrap();
let err = SessionBuilder::new()
.model(dir.path())
.family(*path)
.build()
.unwrap_err();
assert!(
matches!(err, EngineError::Unsupported(_)),
"{path}: {err:?}"
);
continue;
}
assert!(require_stage_b(path).is_ok(), "{path}");
let dir = tempfile::tempdir().unwrap();
write_tiny_q4_bundle(dir.path()).unwrap();
let mut s = SessionBuilder::new()
.model(dir.path())
.family(*path)
.build()
.unwrap();
assert_eq!(s.arch(), *arch);
assert!(!s.graph_hook_name().is_empty());
let gen = s
.generate(
&s.encode_text("ok"),
&GenerateOpts {
max_tokens: 2,
temperature: 0.0,
},
)
.unwrap();
assert!(!gen.tokens.is_empty(), "{path}");
}
}
#[test]
fn stage_c_vl_vla_hooks() {
let dir = tempfile::tempdir().unwrap();
write_tiny_q4_bundle(dir.path()).unwrap();
let s = SessionBuilder::new()
.model(dir.path())
.family("lfm/lfm2-vl-450m")
.build()
.unwrap();
let rgb = vec![10u8; 3 * 4 * 4];
let err = s.vision_prefix(&rgb, 4, 4).unwrap_err();
assert!(matches!(err, EngineError::Unsupported(_)));
let vla = SessionBuilder::new()
.model(dir.path())
.family("openvla/openvla-7b")
.build()
.unwrap();
let err = vla.predict_action("move", 7).unwrap_err();
assert!(matches!(err, EngineError::Unsupported(_)));
let emb = vla.embed_text("hello").unwrap();
assert_eq!(emb.len(), vla.config().hidden_size);
}
#[test]
fn unknown_family() {
let err = SessionBuilder::new()
.model("/tmp")
.family("no/such-model")
.build()
.unwrap_err();
assert!(matches!(err, EngineError::UnsupportedFamily(_)));
}
#[test]
fn greedy_deterministic() {
let dir = tempfile::tempdir().unwrap();
write_tiny_q4_bundle(dir.path()).unwrap();
let mut s = SessionBuilder::new()
.model(dir.path())
.family("gemma/gemma-4-e2b-it")
.build()
.unwrap();
let prompt = s.encode_text("hi");
let opts = GenerateOpts {
max_tokens: 3,
temperature: 0.0,
};
let a = s.generate(&prompt, &opts).unwrap();
let b = s.generate(&prompt, &opts).unwrap();
assert_eq!(a.tokens, b.tokens);
assert_eq!(a.tokens.len(), 3);
}
#[test]
fn encode_chat_is_longer_than_raw_user_text() {
let dir = tempfile::tempdir().unwrap();
write_tiny_q4_bundle(dir.path()).unwrap();
let s = SessionBuilder::new()
.model(dir.path())
.family("qwen/qwen3-0.6b")
.build()
.unwrap();
let raw = s.encode_text("Hello");
let chat = s.encode_chat(&[ChatTurn::new("user", "Hello")]);
assert!(
chat.len() > raw.len(),
"chat template should wrap the user turn (raw={}, chat={})",
raw.len(),
chat.len()
);
assert!(
(s.config().rope_theta - 1_000_000.0).abs() < 1.0,
"Qwen3 must not keep Llama-default rope_theta=10000, got {}",
s.config().rope_theta
);
}
#[test]
fn incremental_decode_matches_full_recompute() {
let dir = tempfile::tempdir().unwrap();
write_tiny_q4_bundle(dir.path()).unwrap();
let mut s = SessionBuilder::new()
.model(dir.path())
.family("gemma/gemma-4-e2b-it")
.build()
.unwrap();
let prompt = s.encode_text("hi");
let max_tokens = 5usize;
let mut prefix = prompt.clone();
if prefix.is_empty() {
prefix.push(1);
}
let mut full_tokens = Vec::new();
for _ in 0..max_tokens {
let logits = s.forward(&prefix).unwrap();
let next = argmax(&logits);
full_tokens.push(next);
prefix.push(next);
if s.is_stop_id(next) {
full_tokens.pop();
break;
}
}
let incr = s
.generate(
&prompt,
&GenerateOpts {
max_tokens,
temperature: 0.0,
},
)
.unwrap();
assert_eq!(
incr.tokens, full_tokens,
"incremental decode must match full-recompute greedy tokens"
);
}
#[test]
fn profile_records_load_and_generate() {
let dir = tempfile::tempdir().unwrap();
write_tiny_q4_bundle(dir.path()).unwrap();
let mut s = SessionBuilder::new()
.model(dir.path())
.family("gemma/gemma-4-e2b-it")
.compute(ComputePref::Cpu)
.profile(true)
.build()
.unwrap();
assert!(s.compute_label().contains("cpu"));
let load = s.last_profile().expect("load profile");
assert!(!load.ci_fail);
assert!(load.load.materialize_ms >= 0.0);
s.generate(
&s.encode_text("hi"),
&GenerateOpts {
max_tokens: 2,
temperature: 0.0,
},
)
.unwrap();
let p = s.last_profile().expect("generate profile");
let g = p.generate.as_ref().expect("generate timings");
assert!(g.prefill_ms >= 0.0);
assert!(g.decode_ms >= 0.0);
}
#[test]
fn cuda_greedy_matches_cpu_if_available() {
if resolve_compute(ComputePref::Cuda).is_err() {
return;
}
let dir = tempfile::tempdir().unwrap();
write_tiny_q4_bundle(dir.path()).unwrap();
let prompt_text = "hi";
let opts = GenerateOpts {
max_tokens: 4,
temperature: 0.0,
};
let mut cpu = SessionBuilder::new()
.model(dir.path())
.family("gemma/gemma-4-e2b-it")
.compute(ComputePref::Cpu)
.build()
.unwrap();
let mut gpu = SessionBuilder::new()
.model(dir.path())
.family("gemma/gemma-4-e2b-it")
.compute(ComputePref::Cuda)
.build()
.unwrap();
assert!(gpu.compute_label().contains("cuda"));
let prompt = cpu.encode_text(prompt_text);
let a = cpu.generate(&prompt, &opts).unwrap();
let b = gpu.generate(&prompt, &opts).unwrap();
assert_eq!(
a.tokens, b.tokens,
"CUDA greedy tokens must match CPU (tiny bundle)"
);
}
#[test]
fn max_tokens_zero_rejected() {
let dir = tempfile::tempdir().unwrap();
write_tiny_q4_bundle(dir.path()).unwrap();
let mut s = SessionBuilder::new()
.model(dir.path())
.family("gemma/gemma-4-e2b-it")
.build()
.unwrap();
let err = s
.generate(
&s.encode_text("x"),
&GenerateOpts {
max_tokens: 0,
temperature: 0.0,
},
)
.unwrap_err();
assert!(matches!(err, EngineError::InvalidParam(_)));
}
#[test]
fn moe_family_refuses_dense_stub() {
let dir = tempfile::tempdir().unwrap();
write_tiny_q4_bundle(dir.path()).unwrap();
let err = SessionBuilder::new()
.model(dir.path())
.family("lfm/lfm2-8b-a1b")
.build()
.unwrap_err();
assert!(matches!(err, EngineError::Unsupported(_)));
assert_eq!(
lookup_family("lfm/lfm2-8b-a1b").unwrap().arch,
ArchClass::TextMoE
);
assert_eq!(graph_hook(ArchClass::TextMoE), "text_moe_decoder");
}
#[test]
fn geometry_gates_conv_and_experts() {
let dir = tempfile::tempdir().unwrap();
write_tiny_q4_bundle(dir.path()).unwrap();
let cfg_path = dir.path().join("config.json");
let mut cfg: Value =
serde_json::from_str(&std::fs::read_to_string(&cfg_path).unwrap()).unwrap();
cfg["model"]["layer_types"] = json!(["conv", "full_attention"]);
std::fs::write(&cfg_path, cfg.to_string()).unwrap();
let err = SessionBuilder::new()
.model(dir.path())
.family("lfm/lfm2-350m")
.build()
.unwrap_err();
assert!(
matches!(err, EngineError::Format(_)),
"expected missing conv tensors, got {err:?}"
);
let dir2 = tempfile::tempdir().unwrap();
write_tiny_q4_bundle(dir2.path()).unwrap();
let cfg_path = dir2.path().join("config.json");
let mut cfg: Value =
serde_json::from_str(&std::fs::read_to_string(&cfg_path).unwrap()).unwrap();
cfg["model"]["layer_types"] = json!(["linear_attention", "full_attention"]);
std::fs::write(&cfg_path, cfg.to_string()).unwrap();
let err = SessionBuilder::new()
.model(dir2.path())
.family("gemma/gemma-3-270m-it")
.build()
.unwrap_err();
assert!(
matches!(err, EngineError::Format(_)),
"expected missing DeltaNet tensors, got {err:?}"
);
}
#[test]
fn lfm_short_conv_and_attn_generate() {
let dir = tempfile::tempdir().unwrap();
let hidden = 8usize;
let layers = 2usize;
let inter = 16usize;
let vocab = 16usize;
let n_heads = 2usize;
let n_kv = 1usize;
let head_dim = 4usize;
let q_dim = n_heads * head_dim;
let k_dim = n_kv * head_dim;
let kernel = 3usize;
let mut tensors = serde_json::Map::new();
let mut bin = Vec::new();
let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
let offset = bin.len();
for &v in data {
bin.extend_from_slice(&v.to_le_bytes());
}
let nbytes = data.len() * 4;
let mut meta = serde_json::Map::new();
meta.insert("kind".into(), json!("raw"));
meta.insert("dtype".into(), json!("f32"));
meta.insert("shape".into(), json!(shape));
meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
tensors.insert(name.to_string(), Value::Object(meta));
};
let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
let n1 = vec![1.0f32; hidden];
add_raw("model.layers.0.operator_norm.weight", vec![hidden], &n1);
add_raw("model.layers.0.ffn_norm.weight", vec![hidden], &n1);
let in_proj = vec![0.02f32; 3 * hidden * hidden];
let out_proj = vec![0.02f32; hidden * hidden];
let conv_w = vec![0.1f32; hidden * kernel];
add_raw(
"model.layers.0.conv.in_proj.weight",
vec![3 * hidden, hidden],
&in_proj,
);
add_raw(
"model.layers.0.conv.out_proj.weight",
vec![hidden, hidden],
&out_proj,
);
add_raw(
"model.layers.0.conv.conv.weight",
vec![hidden, kernel],
&conv_w,
);
let g = vec![0.02f32; inter * hidden];
let d = vec![0.02f32; hidden * inter];
add_raw(
"model.layers.0.mlp.gate_proj.weight",
vec![inter, hidden],
&g,
);
add_raw("model.layers.0.mlp.up_proj.weight", vec![inter, hidden], &g);
add_raw(
"model.layers.0.mlp.down_proj.weight",
vec![hidden, inter],
&d,
);
add_raw("model.layers.1.operator_norm.weight", vec![hidden], &n1);
add_raw(
"model.layers.1.post_attention_layernorm.weight",
vec![hidden],
&n1,
);
let wq = vec![0.02f32; q_dim * hidden];
let wk = vec![0.02f32; k_dim * hidden];
let wv = vec![0.02f32; k_dim * hidden];
let wo = vec![0.02f32; hidden * q_dim];
add_raw(
"model.layers.1.self_attn.q_proj.weight",
vec![q_dim, hidden],
&wq,
);
add_raw(
"model.layers.1.self_attn.k_proj.weight",
vec![k_dim, hidden],
&wk,
);
add_raw(
"model.layers.1.self_attn.v_proj.weight",
vec![k_dim, hidden],
&wv,
);
add_raw(
"model.layers.1.self_attn.o_proj.weight",
vec![hidden, q_dim],
&wo,
);
add_raw(
"model.layers.1.mlp.gate_proj.weight",
vec![inter, hidden],
&g,
);
add_raw("model.layers.1.mlp.up_proj.weight", vec![inter, hidden], &g);
add_raw(
"model.layers.1.mlp.down_proj.weight",
vec![hidden, inter],
&d,
);
add_raw("model.norm.weight", vec![hidden], &n1);
add_raw("lm_head.weight", vec![vocab, hidden], &emb);
let cfg = json!({
"format": "aria-quant-bundle",
"format_version": 2,
"quantization": "test",
"group_size_default": 32,
"hadamard_seed": 0,
"model": {
"hidden_size": hidden,
"num_layers": layers,
"num_attention_heads": n_heads,
"num_kv_heads": n_kv,
"head_dim": head_dim,
"intermediate_size": inter,
"vocab_size": vocab,
"context_length": 32,
"rope_theta": 10000.0,
"conv_l_cache": kernel,
"layer_types": ["conv", "full_attention"]
},
"tensors": tensors
});
std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
let mut s = SessionBuilder::new()
.model(dir.path())
.family("lfm/lfm2-350m")
.build()
.unwrap();
let gen = s
.generate(
&[1, 2, 3],
&GenerateOpts {
max_tokens: 2,
temperature: 0.0,
},
)
.unwrap();
assert_eq!(gen.tokens.len(), 2);
}
#[test]
fn moe_topk_experts_generate() {
let dir = tempfile::tempdir().unwrap();
let hidden = 8usize;
let layers = 1usize;
let inter = 16usize;
let vocab = 16usize;
let n_heads = 2usize;
let n_kv = 1usize;
let head_dim = 4usize;
let q_dim = n_heads * head_dim;
let k_dim = n_kv * head_dim;
let n_experts = 4usize;
let top_k = 2usize;
let mut tensors = serde_json::Map::new();
let mut bin = Vec::new();
let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
let offset = bin.len();
for &v in data {
bin.extend_from_slice(&v.to_le_bytes());
}
let nbytes = data.len() * 4;
let mut meta = serde_json::Map::new();
meta.insert("kind".into(), json!("raw"));
meta.insert("dtype".into(), json!("f32"));
meta.insert("shape".into(), json!(shape));
meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
tensors.insert(name.to_string(), Value::Object(meta));
};
let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
let n1 = vec![1.0f32; hidden];
add_raw("model.layers.0.input_layernorm.weight", vec![hidden], &n1);
add_raw(
"model.layers.0.post_attention_layernorm.weight",
vec![hidden],
&n1,
);
let wq = vec![0.02f32; q_dim * hidden];
let wk = vec![0.02f32; k_dim * hidden];
let wv = vec![0.02f32; k_dim * hidden];
let wo = vec![0.02f32; hidden * q_dim];
add_raw(
"model.layers.0.self_attn.q_proj.weight",
vec![q_dim, hidden],
&wq,
);
add_raw(
"model.layers.0.self_attn.k_proj.weight",
vec![k_dim, hidden],
&wk,
);
add_raw(
"model.layers.0.self_attn.v_proj.weight",
vec![k_dim, hidden],
&wv,
);
add_raw(
"model.layers.0.self_attn.o_proj.weight",
vec![hidden, q_dim],
&wo,
);
let router: Vec<f32> = (0..n_experts * hidden)
.map(|i| ((i % n_experts) as f32) * 0.1)
.collect();
add_raw(
"model.layers.0.block_sparse_moe.gate.weight",
vec![n_experts, hidden],
&router,
);
let g = vec![0.02f32; inter * hidden];
let d = vec![0.02f32; hidden * inter];
for e in 0..n_experts {
add_raw(
&format!("model.layers.0.block_sparse_moe.experts.{e}.w1.weight"),
vec![inter, hidden],
&g,
);
add_raw(
&format!("model.layers.0.block_sparse_moe.experts.{e}.w3.weight"),
vec![inter, hidden],
&g,
);
add_raw(
&format!("model.layers.0.block_sparse_moe.experts.{e}.w2.weight"),
vec![hidden, inter],
&d,
);
}
add_raw("model.norm.weight", vec![hidden], &n1);
add_raw("lm_head.weight", vec![vocab, hidden], &emb);
let cfg = json!({
"format": "aria-quant-bundle",
"format_version": 2,
"quantization": "test",
"group_size_default": 32,
"hadamard_seed": 0,
"model": {
"hidden_size": hidden,
"num_layers": layers,
"num_attention_heads": n_heads,
"num_kv_heads": n_kv,
"head_dim": head_dim,
"intermediate_size": inter,
"vocab_size": vocab,
"context_length": 32,
"rope_theta": 10000.0,
"num_experts": n_experts,
"num_experts_per_tok": top_k
},
"tensors": tensors
});
std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
let mut s = SessionBuilder::new()
.model(dir.path())
.family("inkling/inkling-small")
.build()
.unwrap();
assert_eq!(s.arch(), ArchClass::TextMoE);
assert_eq!(s.graph_hook_name(), "text_moe_decoder");
let gen = s
.generate(
&[1, 2],
&GenerateOpts {
max_tokens: 2,
temperature: 0.0,
},
)
.unwrap();
assert_eq!(gen.tokens.len(), 2);
}
#[test]
fn tiny_q4_codebook_weights_unrotate_on_load() {
let dir = tempfile::tempdir().unwrap();
write_tiny_q4_bundle(dir.path()).unwrap();
let b = load_bundle(dir.path()).unwrap();
let w = b.weight_loaded("blk.0.attn_q.weight").unwrap();
assert!(
w.hdm_seed.is_none(),
"reconstruct_weight path stores original-space W for linear()"
);
let mut s = SessionBuilder::new()
.model(dir.path())
.family("gemma/gemma-4-e2b-it")
.build()
.unwrap();
let gen = s
.generate(
&[1, 2],
&GenerateOpts {
max_tokens: 2,
temperature: 0.0,
},
)
.unwrap();
assert_eq!(gen.tokens.len(), 2);
}
#[test]
fn gemma_hidden_act_geglu_and_qk_norm() {
let dir = tempfile::tempdir().unwrap();
let hidden = 8usize;
let layers = 1usize;
let inter = 16usize;
let vocab = 16usize;
let n_heads = 2usize;
let n_kv = 1usize;
let head_dim = 4usize;
let q_dim = n_heads * head_dim;
let k_dim = n_kv * head_dim;
let p = "model.language_model";
let mut tensors = serde_json::Map::new();
let mut bin = Vec::new();
let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
let offset = bin.len();
for &v in data {
bin.extend_from_slice(&v.to_le_bytes());
}
let nbytes = data.len() * 4;
let mut meta = serde_json::Map::new();
meta.insert("kind".into(), json!("raw"));
meta.insert("dtype".into(), json!("f32"));
meta.insert("shape".into(), json!(shape));
meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
tensors.insert(name.to_string(), Value::Object(meta));
};
let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
add_raw(
&format!("{p}.embed_tokens.weight"),
vec![vocab, hidden],
&emb,
);
let n1 = vec![1.0f32; hidden];
let qn = vec![1.0f32; head_dim];
let kn = vec![1.0f32; head_dim];
let wq = vec![0.01f32; q_dim * hidden];
let wk = vec![0.01f32; k_dim * hidden];
let wv = vec![0.01f32; k_dim * hidden];
let wo = vec![0.01f32; hidden * q_dim];
add_raw(
&format!("{p}.layers.0.input_layernorm.weight"),
vec![hidden],
&n1,
);
add_raw(
&format!("{p}.layers.0.pre_feedforward_layernorm.weight"),
vec![hidden],
&n1,
);
add_raw(
&format!("{p}.layers.0.self_attn.q_proj.weight"),
vec![q_dim, hidden],
&wq,
);
add_raw(
&format!("{p}.layers.0.self_attn.k_proj.weight"),
vec![k_dim, hidden],
&wk,
);
add_raw(
&format!("{p}.layers.0.self_attn.v_proj.weight"),
vec![k_dim, hidden],
&wv,
);
add_raw(
&format!("{p}.layers.0.self_attn.o_proj.weight"),
vec![hidden, q_dim],
&wo,
);
add_raw(
&format!("{p}.layers.0.self_attn.q_norm.weight"),
vec![head_dim],
&qn,
);
add_raw(
&format!("{p}.layers.0.self_attn.k_norm.weight"),
vec![head_dim],
&kn,
);
let g = vec![0.01f32; inter * hidden];
let d = vec![0.01f32; hidden * inter];
add_raw(
&format!("{p}.layers.0.mlp.gate_proj.weight"),
vec![inter, hidden],
&g,
);
add_raw(
&format!("{p}.layers.0.mlp.up_proj.weight"),
vec![inter, hidden],
&g,
);
add_raw(
&format!("{p}.layers.0.mlp.down_proj.weight"),
vec![hidden, inter],
&d,
);
add_raw(&format!("{p}.norm.weight"), vec![hidden], &n1);
add_raw("lm_head.weight", vec![vocab, hidden], &emb);
let cfg = json!({
"format": "aria-quant-bundle",
"format_version": 2,
"quantization": "test",
"group_size_default": 32,
"hadamard_seed": 0,
"model": {
"hidden_size": hidden,
"num_layers": layers,
"num_attention_heads": n_heads,
"num_kv_heads": n_kv,
"head_dim": head_dim,
"global_head_dim": head_dim,
"sliding_window": 512,
"partial_rotary_factor": 0.25,
"intermediate_size": inter,
"vocab_size": vocab,
"context_length": 32,
"rope_theta": 10000.0,
"hidden_act": "gelu_pytorch_tanh",
"layer_types": ["full_attention"]
},
"tensors": tensors
});
std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
let mut s = SessionBuilder::new()
.model(dir.path())
.family("gemma/gemma-4-e2b-it")
.build()
.unwrap();
assert_eq!(s.config().hidden_act.as_deref(), Some("gelu_pytorch_tanh"));
let gen = s
.generate(
&[1, 2],
&GenerateOpts {
max_tokens: 2,
temperature: 0.0,
},
)
.unwrap();
assert_eq!(gen.tokens.len(), 2);
}
#[test]
fn gated_deltanet_and_full_attn_generate() {
let dir = tempfile::tempdir().unwrap();
let hidden = 8usize;
let inter = 16usize;
let vocab = 16usize;
let n_heads = 2usize;
let n_kv = 1usize;
let head_dim = 4usize;
let q_dim = n_heads * head_dim;
let k_dim = n_kv * head_dim;
let n_lin = 2usize;
let hk = 4usize;
let hv = 4usize;
let key_dim = n_lin * hk;
let value_dim = n_lin * hv;
let conv_k = 4usize;
let conv_dim = key_dim * 2 + value_dim;
let mut tensors = serde_json::Map::new();
let mut bin = Vec::new();
let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
let offset = bin.len();
for &v in data {
bin.extend_from_slice(&v.to_le_bytes());
}
let nbytes = data.len() * 4;
let mut meta = serde_json::Map::new();
meta.insert("kind".into(), json!("raw"));
meta.insert("dtype".into(), json!("f32"));
meta.insert("shape".into(), json!(shape));
meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
tensors.insert(name.to_string(), Value::Object(meta));
};
let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
let n1 = vec![1.0f32; hidden];
add_raw("model.layers.0.input_layernorm.weight", vec![hidden], &n1);
add_raw(
"model.layers.0.post_attention_layernorm.weight",
vec![hidden],
&n1,
);
let qkvz = vec![0.02f32; (2 * key_dim + 2 * value_dim) * hidden];
let ba = vec![0.1f32; 2 * n_lin * hidden];
let conv = vec![0.05f32; conv_dim * conv_k];
let a_log = vec![0.5f32; n_lin];
let dt = vec![1.0f32; n_lin];
let outp = vec![0.02f32; hidden * value_dim];
add_raw(
"model.layers.0.linear_attn.in_proj_qkvz.weight",
vec![2 * key_dim + 2 * value_dim, hidden],
&qkvz,
);
add_raw(
"model.layers.0.linear_attn.in_proj_ba.weight",
vec![2 * n_lin, hidden],
&ba,
);
add_raw(
"model.layers.0.linear_attn.conv1d.weight",
vec![conv_dim, conv_k],
&conv,
);
add_raw("model.layers.0.linear_attn.A_log", vec![n_lin], &a_log);
add_raw("model.layers.0.linear_attn.dt_bias", vec![n_lin], &dt);
add_raw(
"model.layers.0.linear_attn.out_proj.weight",
vec![hidden, value_dim],
&outp,
);
let g = vec![0.02f32; inter * hidden];
let d = vec![0.02f32; hidden * inter];
add_raw(
"model.layers.0.mlp.gate_proj.weight",
vec![inter, hidden],
&g,
);
add_raw("model.layers.0.mlp.up_proj.weight", vec![inter, hidden], &g);
add_raw(
"model.layers.0.mlp.down_proj.weight",
vec![hidden, inter],
&d,
);
add_raw("model.layers.1.input_layernorm.weight", vec![hidden], &n1);
add_raw(
"model.layers.1.post_attention_layernorm.weight",
vec![hidden],
&n1,
);
let wq = vec![0.02f32; q_dim * hidden];
let wk = vec![0.02f32; k_dim * hidden];
let wv = vec![0.02f32; k_dim * hidden];
let wo = vec![0.02f32; hidden * q_dim];
add_raw(
"model.layers.1.self_attn.q_proj.weight",
vec![q_dim, hidden],
&wq,
);
add_raw(
"model.layers.1.self_attn.k_proj.weight",
vec![k_dim, hidden],
&wk,
);
add_raw(
"model.layers.1.self_attn.v_proj.weight",
vec![k_dim, hidden],
&wv,
);
add_raw(
"model.layers.1.self_attn.o_proj.weight",
vec![hidden, q_dim],
&wo,
);
add_raw(
"model.layers.1.mlp.gate_proj.weight",
vec![inter, hidden],
&g,
);
add_raw("model.layers.1.mlp.up_proj.weight", vec![inter, hidden], &g);
add_raw(
"model.layers.1.mlp.down_proj.weight",
vec![hidden, inter],
&d,
);
add_raw("model.norm.weight", vec![hidden], &n1);
add_raw("lm_head.weight", vec![vocab, hidden], &emb);
let cfg = json!({
"format": "aria-quant-bundle",
"format_version": 2,
"quantization": "test",
"hadamard_seed": 0,
"model": {
"hidden_size": hidden,
"num_layers": 2,
"num_attention_heads": n_heads,
"num_kv_heads": n_kv,
"head_dim": head_dim,
"intermediate_size": inter,
"vocab_size": vocab,
"context_length": 32,
"rope_theta": 10000.0,
"layer_types": ["linear_attention", "full_attention"]
},
"tensors": tensors
});
std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
let mut s = SessionBuilder::new()
.model(dir.path())
.family("qwen/qwen3.5-2b")
.build()
.unwrap();
assert_eq!(s.config().partial_rotary_factor, Some(0.25));
assert!(
(s.config().rope_theta - 10_000_000.0).abs() < 1.0,
"Qwen3.5 Llama-default rope_theta must become 1e7, got {}",
s.config().rope_theta
);
let gen = s
.generate(
&[1, 2, 3],
&GenerateOpts {
max_tokens: 2,
temperature: 0.0,
},
)
.unwrap();
assert_eq!(gen.tokens.len(), 2);
}
#[test]
fn gated_deltanet_split_qwen35_projections_generate() {
let dir = tempfile::tempdir().unwrap();
let hidden = 8usize;
let inter = 16usize;
let vocab = 16usize;
let n_heads = 2usize;
let n_kv = 1usize;
let head_dim = 4usize;
let q_dim = n_heads * head_dim;
let k_dim = n_kv * head_dim;
let n_lin = 2usize;
let hk = 4usize;
let hv = 4usize;
let key_dim = n_lin * hk;
let value_dim = n_lin * hv;
let conv_k = 4usize;
let conv_dim = key_dim * 2 + value_dim;
let mut tensors = serde_json::Map::new();
let mut bin = Vec::new();
let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
let offset = bin.len();
for &v in data {
bin.extend_from_slice(&v.to_le_bytes());
}
let nbytes = data.len() * 4;
let mut meta = serde_json::Map::new();
meta.insert("kind".into(), json!("raw"));
meta.insert("dtype".into(), json!("f32"));
meta.insert("shape".into(), json!(shape));
meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
tensors.insert(name.to_string(), Value::Object(meta));
};
let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
add_raw("model.embed_tokens.weight", vec![vocab, hidden], &emb);
let n1 = vec![1.0f32; hidden];
add_raw("model.layers.0.input_layernorm.weight", vec![hidden], &n1);
add_raw(
"model.layers.0.post_attention_layernorm.weight",
vec![hidden],
&n1,
);
let qkv = vec![0.02f32; (2 * key_dim + value_dim) * hidden];
let z = vec![0.02f32; value_dim * hidden];
let proj_b = vec![0.1f32; n_lin * hidden];
let proj_a = vec![0.1f32; n_lin * hidden];
let conv = vec![0.05f32; conv_dim * conv_k];
let a_log = vec![0.5f32; n_lin];
let dt = vec![1.0f32; n_lin];
let outp = vec![0.02f32; hidden * value_dim];
add_raw(
"model.layers.0.linear_attn.in_proj_qkv.weight",
vec![2 * key_dim + value_dim, hidden],
&qkv,
);
add_raw(
"model.layers.0.linear_attn.in_proj_z.weight",
vec![value_dim, hidden],
&z,
);
add_raw(
"model.layers.0.linear_attn.in_proj_b.weight",
vec![n_lin, hidden],
&proj_b,
);
add_raw(
"model.layers.0.linear_attn.in_proj_a.weight",
vec![n_lin, hidden],
&proj_a,
);
add_raw(
"model.layers.0.linear_attn.conv1d.weight",
vec![conv_dim, conv_k],
&conv,
);
add_raw("model.layers.0.linear_attn.A_log", vec![n_lin], &a_log);
add_raw("model.layers.0.linear_attn.dt_bias", vec![n_lin], &dt);
add_raw(
"model.layers.0.linear_attn.out_proj.weight",
vec![hidden, value_dim],
&outp,
);
let g = vec![0.02f32; inter * hidden];
let d = vec![0.02f32; hidden * inter];
add_raw(
"model.layers.0.mlp.gate_proj.weight",
vec![inter, hidden],
&g,
);
add_raw("model.layers.0.mlp.up_proj.weight", vec![inter, hidden], &g);
add_raw(
"model.layers.0.mlp.down_proj.weight",
vec![hidden, inter],
&d,
);
add_raw("model.layers.1.input_layernorm.weight", vec![hidden], &n1);
add_raw(
"model.layers.1.post_attention_layernorm.weight",
vec![hidden],
&n1,
);
let wq = vec![0.02f32; q_dim * hidden];
let wk = vec![0.02f32; k_dim * hidden];
let wv = vec![0.02f32; k_dim * hidden];
let wo = vec![0.02f32; hidden * q_dim];
add_raw(
"model.layers.1.self_attn.q_proj.weight",
vec![q_dim, hidden],
&wq,
);
add_raw(
"model.layers.1.self_attn.k_proj.weight",
vec![k_dim, hidden],
&wk,
);
add_raw(
"model.layers.1.self_attn.v_proj.weight",
vec![k_dim, hidden],
&wv,
);
add_raw(
"model.layers.1.self_attn.o_proj.weight",
vec![hidden, q_dim],
&wo,
);
add_raw(
"model.layers.1.mlp.gate_proj.weight",
vec![inter, hidden],
&g,
);
add_raw("model.layers.1.mlp.up_proj.weight", vec![inter, hidden], &g);
add_raw(
"model.layers.1.mlp.down_proj.weight",
vec![hidden, inter],
&d,
);
add_raw("model.norm.weight", vec![hidden], &n1);
add_raw("lm_head.weight", vec![vocab, hidden], &emb);
let cfg = json!({
"format": "aria-quant-bundle",
"format_version": 2,
"quantization": "test",
"hadamard_seed": 0,
"model": {
"hidden_size": hidden,
"num_layers": 2,
"num_attention_heads": n_heads,
"num_kv_heads": n_kv,
"head_dim": head_dim,
"intermediate_size": inter,
"vocab_size": vocab,
"context_length": 32,
"rope_theta": 10000.0,
"layer_types": ["linear_attention", "full_attention"]
},
"tensors": tensors
});
std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
let mut s = SessionBuilder::new()
.model(dir.path())
.family("qwen/qwen3.5-0.8b")
.build()
.unwrap();
let gen = s
.generate(
&[1, 2, 3],
&GenerateOpts {
max_tokens: 2,
temperature: 0.0,
},
)
.unwrap();
assert_eq!(gen.tokens.len(), 2);
}
#[test]
fn vision_and_action_consume_bundle_weights() {
let dir = tempfile::tempdir().unwrap();
write_tiny_q4_bundle(dir.path()).unwrap();
let cfg_path = dir.path().join("config.json");
let mut cfg: Value =
serde_json::from_str(&std::fs::read_to_string(&cfg_path).unwrap()).unwrap();
let hidden = cfg["model"]["hidden_size"].as_u64().unwrap() as usize;
let mut tensors = cfg["tensors"].as_object().cloned().unwrap();
let mut bin = std::fs::read(dir.path().join("weight.bin")).unwrap();
let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
let offset = bin.len();
for &v in data {
bin.extend_from_slice(&v.to_le_bytes());
}
let nbytes = data.len() * 4;
tensors.insert(
name.to_string(),
json!({
"kind": "raw",
"dtype": "f32",
"shape": shape,
"offsets": { "data": [offset, nbytes] }
}),
);
};
let vis = vec![0.1f32; hidden * 3];
add_raw("mm_projector.weight", vec![hidden, 3], &vis);
let act_dim = 7usize;
let act = vec![0.05f32; act_dim * hidden];
add_raw("action_head.weight", vec![act_dim, hidden], &act);
cfg["tensors"] = Value::Object(tensors);
std::fs::write(&cfg_path, cfg.to_string()).unwrap();
std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
let s = SessionBuilder::new()
.model(dir.path())
.family("lfm/lfm2-vl-450m")
.build()
.unwrap();
let rgb = vec![10u8; 3 * 4 * 4];
let pref = s.vision_prefix(&rgb, 4, 4).unwrap();
assert_eq!(pref.len(), hidden);
let vla = SessionBuilder::new()
.model(dir.path())
.family("openvla/openvla-7b")
.build()
.unwrap();
let a = vla.predict_action("move", act_dim).unwrap();
assert_eq!(a.len(), act_dim);
}
#[test]
fn load_real_hf_named_bundle_if_present() {
let Ok(path) = std::env::var("ARIA_SMOKE_BUNDLE") else {
return;
};
let path = std::path::Path::new(&path);
if !path.join("config.json").is_file() {
return;
}
let family = if path.to_string_lossy().contains("gemma-4") {
"gemma/gemma-4-e2b-it"
} else {
"qwen/qwen3-0.6b"
};
let s = SessionBuilder::new()
.model(path)
.family(family)
.build()
.unwrap_or_else(|e| panic!("{family} bundle should materialize: {e}"));
assert!(s.config().num_layers > 0);
assert!(s.config().hidden_size > 0);
if family.contains("gemma-4") && s.config().hidden_size >= 1024 {
assert!(
s.weights.ple.is_some(),
"real Gemma-4 q4 must load codebook PLE"
);
let hidden = s.config().hidden_size;
let vocab = s.config().vocab_size;
assert!(
s.weights.emb.data.len() >= vocab.saturating_mul(hidden),
"embed table too small for vocab={vocab} hidden={hidden}"
);
}
}
#[test]
fn gemma4_four_norm_ple_and_tied_embed_generate() {
let dir = tempfile::tempdir().unwrap();
let hidden = 8usize;
let layers = 1usize;
let inter = 16usize;
let vocab = 16usize;
let n_heads = 2usize;
let n_kv = 1usize;
let head_dim = 4usize;
let q_dim = n_heads * head_dim;
let k_dim = n_kv * head_dim;
let ple_d = 4usize;
let p = "model.language_model";
let mut tensors = serde_json::Map::new();
let mut bin = Vec::new();
let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
let offset = bin.len();
for &v in data {
bin.extend_from_slice(&v.to_le_bytes());
}
let nbytes = data.len() * 4;
let mut meta = serde_json::Map::new();
meta.insert("kind".into(), json!("raw"));
meta.insert("dtype".into(), json!("f32"));
meta.insert("shape".into(), json!(shape));
meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
tensors.insert(name.to_string(), Value::Object(meta));
};
let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
add_raw(
&format!("{p}.embed_tokens.weight"),
vec![vocab, hidden],
&emb,
);
let n1 = vec![1.0f32; hidden];
add_raw(
&format!("{p}.layers.0.input_layernorm.weight"),
vec![hidden],
&n1,
);
add_raw(
&format!("{p}.layers.0.post_attention_layernorm.weight"),
vec![hidden],
&n1,
);
add_raw(
&format!("{p}.layers.0.pre_feedforward_layernorm.weight"),
vec![hidden],
&n1,
);
add_raw(
&format!("{p}.layers.0.post_feedforward_layernorm.weight"),
vec![hidden],
&n1,
);
add_raw(&format!("{p}.layers.0.layer_scalar"), vec![1], &[0.5f32]);
let wq = vec![0.01f32; q_dim * hidden];
let wk = vec![0.01f32; k_dim * hidden];
let wv = vec![0.01f32; k_dim * hidden];
let wo = vec![0.01f32; hidden * q_dim];
add_raw(
&format!("{p}.layers.0.self_attn.q_proj.weight"),
vec![q_dim, hidden],
&wq,
);
add_raw(
&format!("{p}.layers.0.self_attn.k_proj.weight"),
vec![k_dim, hidden],
&wk,
);
add_raw(
&format!("{p}.layers.0.self_attn.v_proj.weight"),
vec![k_dim, hidden],
&wv,
);
add_raw(
&format!("{p}.layers.0.self_attn.o_proj.weight"),
vec![hidden, q_dim],
&wo,
);
let g = vec![0.01f32; inter * hidden];
let d = vec![0.01f32; hidden * inter];
add_raw(
&format!("{p}.layers.0.mlp.gate_proj.weight"),
vec![inter, hidden],
&g,
);
add_raw(
&format!("{p}.layers.0.mlp.up_proj.weight"),
vec![inter, hidden],
&g,
);
add_raw(
&format!("{p}.layers.0.mlp.down_proj.weight"),
vec![hidden, inter],
&d,
);
let packed = layers * ple_d;
let ple_emb = vec![0.02f32; vocab * packed];
add_raw(
&format!("{p}.embed_tokens_per_layer.weight"),
vec![vocab, packed],
&ple_emb,
);
let ple_proj = vec![0.01f32; packed * hidden];
add_raw(
&format!("{p}.per_layer_model_projection.weight"),
vec![packed, hidden],
&ple_proj,
);
let ple_pn = vec![1.0f32; ple_d];
add_raw(
&format!("{p}.per_layer_projection_norm.weight"),
vec![ple_d],
&ple_pn,
);
let ple_gate = vec![0.01f32; ple_d * hidden];
let ple_out = vec![0.01f32; hidden * ple_d];
add_raw(
&format!("{p}.layers.0.per_layer_input_gate.weight"),
vec![ple_d, hidden],
&ple_gate,
);
add_raw(
&format!("{p}.layers.0.per_layer_projection.weight"),
vec![hidden, ple_d],
&ple_out,
);
add_raw(
&format!("{p}.layers.0.post_per_layer_input_norm.weight"),
vec![hidden],
&n1,
);
add_raw(&format!("{p}.norm.weight"), vec![hidden], &n1);
add_raw("lm_head.weight", vec![vocab, hidden], &emb);
let cfg = json!({
"format": "aria-quant-bundle",
"format_version": 2,
"quantization": "test",
"group_size_default": 32,
"hadamard_seed": 0,
"model": {
"hidden_size": hidden,
"num_layers": layers,
"num_attention_heads": n_heads,
"num_kv_heads": n_kv,
"intermediate_size": inter,
"vocab_size": vocab,
"context_length": 32,
"rope_theta": 10000.0,
"hidden_act": "gelu_pytorch_tanh",
"tie_word_embeddings": true,
"head_dim": head_dim,
"global_head_dim": head_dim,
"sliding_window": 512,
"partial_rotary_factor": 0.25,
"layer_types": ["full_attention"]
},
"tensors": tensors
});
std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
let mut s = SessionBuilder::new()
.model(dir.path())
.family("gemma/gemma-4-e2b-it")
.build()
.unwrap();
assert!((s.embed_scale - (hidden as f32).sqrt()).abs() < 1e-5);
assert!(s.weights.ple.is_some());
assert!((s.weights.layers[0].layer_scalar - 0.5).abs() < 1e-6);
assert!(s.weights.layers[0].post_attn_norm.is_some());
assert!(s.weights.layers[0].post_ffn_norm.is_some());
let prompt = vec![1u32, 2];
let batched = s
.generate(
&prompt,
&GenerateOpts {
max_tokens: 3,
temperature: 0.0,
},
)
.unwrap();
let step = s
.generate(
&prompt,
&GenerateOpts {
max_tokens: 3,
temperature: 0.0,
},
)
.unwrap();
assert_eq!(batched.tokens, step.tokens);
assert_eq!(batched.tokens.len(), 3);
assert_eq!(s.config().sliding_window, Some(512));
}
#[test]
fn gemma4_ple_required_gate() {
assert!(!gemma4_requires_ple("gemma/gemma-4-e2b-it", 64));
assert!(gemma4_requires_ple("gemma/gemma-4-e2b-it", 1024));
assert!(gemma4_requires_ple("gemma/gemma-4-e2b-it", 1536));
assert!(gemma4_requires_ple("gemma/gemma-3n-e2b-it", 2048));
assert!(!gemma4_requires_ple("qwen/qwen3-0.6b", 1536));
assert!(!gemma3n_requires_altup("gemma/gemma-3n-e2b-it", 64));
assert!(gemma3n_requires_altup("gemma/gemma-3n-e2b-it", 2048));
assert!(!gemma3n_requires_altup("gemma/gemma-4-e2b-it", 2048));
}
#[test]
fn gemma3n_altup_laurel_ple_generate() {
let dir = tempfile::tempdir().unwrap();
let hidden = 8usize;
let layers = 1usize;
let inter = 16usize;
let vocab = 16usize;
let n_heads = 2usize;
let n_kv = 1usize;
let head_dim = 4usize;
let q_dim = n_heads * head_dim;
let k_dim = n_kv * head_dim;
let ple_d = 4usize;
let rank = 2usize;
let n_alt = 4usize;
let p = "model.language_model";
let mut tensors = serde_json::Map::new();
let mut bin = Vec::new();
let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
let offset = bin.len();
for &v in data {
bin.extend_from_slice(&v.to_le_bytes());
}
let nbytes = data.len() * 4;
let mut meta = serde_json::Map::new();
meta.insert("kind".into(), json!("raw"));
meta.insert("dtype".into(), json!("f32"));
meta.insert("shape".into(), json!(shape));
meta.insert("offsets".into(), json!({ "data": [offset, nbytes] }));
tensors.insert(name.to_string(), Value::Object(meta));
};
let emb: Vec<f32> = (0..vocab * hidden).map(|i| i as f32 * 0.01).collect();
add_raw(
&format!("{p}.embed_tokens.weight"),
vec![vocab, hidden],
&emb,
);
let n1 = vec![1.0f32; hidden];
add_raw(
&format!("{p}.layers.0.input_layernorm.weight"),
vec![hidden],
&n1,
);
add_raw(
&format!("{p}.layers.0.post_attention_layernorm.weight"),
vec![hidden],
&n1,
);
add_raw(
&format!("{p}.layers.0.pre_feedforward_layernorm.weight"),
vec![hidden],
&n1,
);
add_raw(
&format!("{p}.layers.0.post_feedforward_layernorm.weight"),
vec![hidden],
&n1,
);
let wq = vec![0.01f32; q_dim * hidden];
let wk = vec![0.01f32; k_dim * hidden];
let wv = vec![0.01f32; k_dim * hidden];
let wo = vec![0.01f32; hidden * q_dim];
add_raw(
&format!("{p}.layers.0.self_attn.q_proj.weight"),
vec![q_dim, hidden],
&wq,
);
add_raw(
&format!("{p}.layers.0.self_attn.k_proj.weight"),
vec![k_dim, hidden],
&wk,
);
add_raw(
&format!("{p}.layers.0.self_attn.v_proj.weight"),
vec![k_dim, hidden],
&wv,
);
add_raw(
&format!("{p}.layers.0.self_attn.o_proj.weight"),
vec![hidden, q_dim],
&wo,
);
let g = vec![0.01f32; inter * hidden];
let d = vec![0.01f32; hidden * inter];
add_raw(
&format!("{p}.layers.0.mlp.gate_proj.weight"),
vec![inter, hidden],
&g,
);
add_raw(
&format!("{p}.layers.0.mlp.up_proj.weight"),
vec![inter, hidden],
&g,
);
add_raw(
&format!("{p}.layers.0.mlp.down_proj.weight"),
vec![hidden, inter],
&d,
);
let packed = layers * ple_d;
add_raw(
&format!("{p}.embed_tokens_per_layer.weight"),
vec![vocab, packed],
&vec![0.02f32; vocab * packed],
);
add_raw(
&format!("{p}.per_layer_model_projection.weight"),
vec![packed, hidden],
&vec![0.01f32; packed * hidden],
);
add_raw(
&format!("{p}.per_layer_projection_norm.weight"),
vec![ple_d],
&vec![1.0f32; ple_d],
);
add_raw(
&format!("{p}.layers.0.per_layer_input_gate.weight"),
vec![ple_d, hidden],
&vec![0.01f32; ple_d * hidden],
);
add_raw(
&format!("{p}.layers.0.per_layer_projection.weight"),
vec![hidden, ple_d],
&vec![0.01f32; hidden * ple_d],
);
add_raw(
&format!("{p}.layers.0.post_per_layer_input_norm.weight"),
vec![hidden],
&n1,
);
add_raw(
&format!("{p}.layers.0.altup.modality_router.weight"),
vec![n_alt, hidden],
&vec![0.01f32; n_alt * hidden],
);
add_raw(
&format!("{p}.layers.0.altup.router_norm.weight"),
vec![hidden],
&n1,
);
add_raw(
&format!("{p}.layers.0.altup.prediction_coefs.weight"),
vec![n_alt * n_alt, n_alt],
&vec![0.0f32; n_alt * n_alt * n_alt],
);
add_raw(
&format!("{p}.layers.0.altup.correction_coefs.weight"),
vec![n_alt, n_alt],
&vec![0.0f32; n_alt * n_alt],
);
add_raw(
&format!("{p}.layers.0.altup.correct_output_scale"),
vec![hidden],
&n1,
);
add_raw(
&format!("{p}.layers.0.laurel.linear_left.weight"),
vec![rank, hidden],
&vec![0.01f32; rank * hidden],
);
add_raw(
&format!("{p}.layers.0.laurel.linear_right.weight"),
vec![hidden, rank],
&vec![0.01f32; hidden * rank],
);
add_raw(
&format!("{p}.layers.0.laurel.post_laurel_norm.weight"),
vec![hidden],
&n1,
);
let eye: Vec<f32> = (0..hidden * hidden)
.map(|i| if i / hidden == i % hidden { 0.05 } else { 0.0 })
.collect();
for i in 0..3 {
add_raw(
&format!("{p}.altup_projections.{i}.weight"),
vec![hidden, hidden],
&eye,
);
add_raw(
&format!("{p}.altup_unembed_projections.{i}.weight"),
vec![hidden, hidden],
&eye,
);
}
add_raw(&format!("{p}.norm.weight"), vec![hidden], &n1);
add_raw("lm_head.weight", vec![vocab, hidden], &emb);
let cfg = json!({
"format": "aria-quant-bundle",
"format_version": 2,
"quantization": "test",
"group_size_default": 32,
"hadamard_seed": 0,
"model": {
"hidden_size": hidden,
"num_layers": layers,
"num_attention_heads": n_heads,
"num_kv_heads": n_kv,
"intermediate_size": inter,
"vocab_size": vocab,
"context_length": 32,
"rope_theta": 1000000.0,
"hidden_act": "gelu_pytorch_tanh",
"tie_word_embeddings": true,
"head_dim": head_dim,
"sliding_window": 512,
"layer_types": ["full_attention"]
},
"tensors": tensors
});
std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
let mut s = SessionBuilder::new()
.model(dir.path())
.family("gemma/gemma-3n-e2b-it")
.build()
.unwrap();
assert!(s.has_gemma3n_graph());
assert!(s.weights.ple.is_some());
assert!(s.weights.layers[0].altup.is_some());
assert!(s.weights.layers[0].laurel.is_some());
let prompt = vec![1u32, 2];
let batched = s
.generate(
&prompt,
&GenerateOpts {
max_tokens: 3,
temperature: 0.0,
},
)
.unwrap();
let step = s
.generate(
&prompt,
&GenerateOpts {
max_tokens: 3,
temperature: 0.0,
},
)
.unwrap();
assert_eq!(batched.tokens, step.tokens);
assert_eq!(batched.tokens.len(), 3);
}
#[test]
fn gaussian_topk_sparsity_zeros_below_cutoff() {
let row: Vec<f32> = (0..100).map(|i| i as f32).collect();
let y = gaussian_topk(&row, 100, 0.95).unwrap();
assert!(y.iter().all(|&v| v >= 0.0));
assert!(y.iter().filter(|&&v| v == 0.0).count() > 50);
assert!(y.iter().any(|&v| v > 0.0));
}
#[test]
fn gemma4_e2b_scale_missing_ple_is_hard_error() {
let dir = tempfile::tempdir().unwrap();
let hidden = 1024usize;
let layers = 1usize;
let vocab = 8usize;
let n_heads = 8usize;
let n_kv = 1usize;
let head_dim = 128usize;
let q_dim = n_heads * head_dim;
let k_dim = n_kv * head_dim;
let inter = 32usize;
let p = "model.language_model";
let mut tensors = serde_json::Map::new();
let mut bin = Vec::new();
let mut add_raw = |name: &str, shape: Vec<usize>, data: &[f32]| {
let offset = bin.len();
for &v in data {
bin.extend_from_slice(&v.to_le_bytes());
}
let nbytes = data.len() * 4;
tensors.insert(
name.to_string(),
json!({
"kind": "raw",
"dtype": "f32",
"shape": shape,
"offsets": { "data": [offset, nbytes] }
}),
);
};
let emb = vec![0.01f32; vocab * hidden];
add_raw(
&format!("{p}.embed_tokens.weight"),
vec![vocab, hidden],
&emb,
);
let ones = vec![1.0f32; hidden];
add_raw(
&format!("{p}.layers.0.input_layernorm.weight"),
vec![hidden],
&ones,
);
add_raw(
&format!("{p}.layers.0.pre_feedforward_layernorm.weight"),
vec![hidden],
&ones,
);
add_raw(
&format!("{p}.layers.0.post_attention_layernorm.weight"),
vec![hidden],
&ones,
);
add_raw(
&format!("{p}.layers.0.post_feedforward_layernorm.weight"),
vec![hidden],
&ones,
);
add_raw(&format!("{p}.norm.weight"), vec![hidden], &ones);
let q = vec![0.01f32; q_dim * hidden];
let k = vec![0.01f32; k_dim * hidden];
add_raw(
&format!("{p}.layers.0.self_attn.q_proj.weight"),
vec![q_dim, hidden],
&q,
);
add_raw(
&format!("{p}.layers.0.self_attn.k_proj.weight"),
vec![k_dim, hidden],
&k,
);
add_raw(
&format!("{p}.layers.0.self_attn.v_proj.weight"),
vec![k_dim, hidden],
&k,
);
add_raw(
&format!("{p}.layers.0.self_attn.o_proj.weight"),
vec![hidden, q_dim],
&q,
);
let qn = vec![1.0f32; head_dim];
add_raw(
&format!("{p}.layers.0.self_attn.q_norm.weight"),
vec![head_dim],
&qn,
);
add_raw(
&format!("{p}.layers.0.self_attn.k_norm.weight"),
vec![head_dim],
&qn,
);
let g = vec![0.01f32; inter * hidden];
add_raw(
&format!("{p}.layers.0.mlp.gate_proj.weight"),
vec![inter, hidden],
&g,
);
add_raw(
&format!("{p}.layers.0.mlp.up_proj.weight"),
vec![inter, hidden],
&g,
);
add_raw(
&format!("{p}.layers.0.mlp.down_proj.weight"),
vec![hidden, inter],
&g,
);
let cfg = json!({
"format": "aria-quant-bundle",
"format_version": 2,
"quantization": "test",
"group_size_default": 32,
"hadamard_seed": 0,
"model": {
"hidden_size": hidden,
"num_layers": layers,
"num_attention_heads": n_heads,
"num_kv_heads": n_kv,
"intermediate_size": inter,
"vocab_size": vocab,
"context_length": 32,
"rope_theta": 10000.0,
"hidden_act": "gelu_pytorch_tanh",
"tie_word_embeddings": true,
"head_dim": head_dim,
"global_head_dim": head_dim,
"sliding_window": 512,
"partial_rotary_factor": 0.25,
"num_kv_shared_layers": 0,
"layer_types": ["full_attention"]
},
"tensors": tensors
});
std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
let err = SessionBuilder::new()
.model(dir.path())
.family("gemma/gemma-4-e2b-it")
.build()
.unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("PLE") && msg.contains("embed_tokens_per_layer"),
"{msg}"
);
}
#[test]
fn gemma4_sliding_window_config_and_generate() {
let dir_wide = tempfile::tempdir().unwrap();
write_tiny_q4_bundle(dir_wide.path()).unwrap();
let dir_narrow = tempfile::tempdir().unwrap();
write_tiny_q4_bundle(dir_narrow.path()).unwrap();
let patch = |path: &std::path::Path, window: usize| {
let cfg_path = path.join("config.json");
let raw = std::fs::read_to_string(&cfg_path).unwrap();
let mut cfg: Value = serde_json::from_str(&raw).unwrap();
cfg["model"]["sliding_window"] = json!(window);
cfg["model"]["layer_types"] = json!(["sliding_attention", "sliding_attention"]);
std::fs::write(&cfg_path, cfg.to_string()).unwrap();
};
patch(dir_wide.path(), 512);
patch(dir_narrow.path(), 1);
let wide = SessionBuilder::new()
.model(dir_wide.path())
.family("gemma/gemma-4-e2b-it")
.build()
.unwrap();
let mut narrow = SessionBuilder::new()
.model(dir_narrow.path())
.family("gemma/gemma-4-e2b-it")
.build()
.unwrap();
assert_eq!(wide.config().sliding_window, Some(512));
assert_eq!(narrow.config().sliding_window, Some(1));
assert_eq!(wide.attn_window(AttnKind::Sliding), Some(512));
assert_eq!(narrow.attn_window(AttnKind::Sliding), Some(1));
for layer in &narrow.weights.layers {
if let LayerOp::Attn(attn) = &layer.op {
assert_eq!(attn.kind, AttnKind::Sliding);
}
}
let prompt = vec![1u32, 2, 3, 4];
let gen = narrow
.generate(
&prompt,
&GenerateOpts {
max_tokens: 3,
temperature: 0.0,
},
)
.unwrap();
assert_eq!(gen.tokens.len(), 3);
let mut incr = SessionBuilder::new()
.model(dir_narrow.path())
.family("gemma/gemma-4-e2b-it")
.build()
.unwrap();
let again = incr
.generate(
&prompt,
&GenerateOpts {
max_tokens: 3,
temperature: 0.0,
},
)
.unwrap();
assert_eq!(gen.tokens, again.tokens);
}
}