Skip to main content

aria_inference/
bundle.rs

1use crate::pack::unpack_indices;
2use aria_kernel::{
3    dequant_lookup_group, hadamard_blocked_rows_tiles, pow2_tile_sizes, EngineError,
4};
5use half::f16;
6use memmap2::Mmap;
7use serde::Deserialize;
8use serde_json::Value;
9use std::collections::HashMap;
10use std::fs::File;
11use std::path::{Path, PathBuf};
12use std::sync::Arc;
13
14pub const BUNDLE_FORMAT: &str = "aria-quant-bundle";
15
16#[derive(Debug, Clone, Deserialize)]
17pub struct ModelConfig {
18    pub hidden_size: usize,
19    pub num_layers: usize,
20    pub num_attention_heads: usize,
21    pub num_kv_heads: usize,
22    pub intermediate_size: usize,
23    pub vocab_size: usize,
24    pub context_length: usize,
25    #[serde(default = "default_rope")]
26    pub rope_theta: f32,
27    #[serde(default)]
28    pub head_dim: Option<usize>,
29    #[serde(default)]
30    pub layer_types: Option<Vec<String>>,
31    #[serde(default)]
32    pub num_kv_shared_layers: Option<usize>,
33    #[serde(default)]
34    pub use_double_wide_mlp: Option<bool>,
35    #[serde(default)]
36    pub hidden_act: Option<String>,
37    #[serde(default)]
38    pub num_experts: Option<usize>,
39    #[serde(default)]
40    pub num_experts_per_tok: Option<usize>,
41    #[serde(default)]
42    pub tie_word_embeddings: Option<bool>,
43    /// LFM2 short-conv kernel / cache length (HF `conv_L_cache`, default 3).
44    #[serde(default)]
45    pub conv_l_cache: Option<usize>,
46    /// Sliding-window attention length (HF `sliding_window`). Required for gemma-4.
47    #[serde(default)]
48    pub sliding_window: Option<usize>,
49    /// Gemma-4 global p-RoPE fraction (HF `partial_rotary_factor`). Required for gemma-4.
50    #[serde(default)]
51    pub partial_rotary_factor: Option<f32>,
52    /// Gemma-4 full-attention head dim (HF `global_head_dim`). Required for gemma-4.
53    #[serde(default)]
54    pub global_head_dim: Option<usize>,
55}
56
57fn default_rope() -> f32 {
58    10000.0
59}
60
61#[derive(Debug, Deserialize)]
62struct BundleConfig {
63    format: String,
64    format_version: u32,
65    #[allow(dead_code)]
66    quantization: String,
67    #[serde(default)]
68    group_size_default: usize,
69    #[serde(default)]
70    hadamard_seed: Option<i64>,
71    model: ModelConfig,
72    tensors: HashMap<String, Value>,
73}
74
75#[derive(Debug, Clone)]
76pub struct QuantTensor {
77    pub bits: u8,
78    pub group_size: usize,
79    pub shape: (usize, usize),
80    pub row_pad: usize,
81    pub codebook_share: String,
82    pub packed_indices: Vec<u8>,
83    pub codebook: Vec<f32>,
84    pub codebook_shape: Vec<usize>,
85    pub hadamard: Value,
86}
87
88#[derive(Debug, Clone)]
89pub enum TensorData {
90    Codebook(QuantTensor),
91    Raw {
92        dtype: String,
93        shape: Vec<usize>,
94        data: Vec<f32>,
95    },
96}
97
98pub struct Bundle {
99    pub path: PathBuf,
100    pub model: ModelConfig,
101    pub quantization: String,
102    pub group_size_default: usize,
103    pub hadamard_seed: Option<i64>,
104    pub tensors: HashMap<String, TensorData>,
105    #[allow(dead_code)]
106    mmap: Arc<Mmap>,
107}
108
109impl std::fmt::Debug for Bundle {
110    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
111        f.debug_struct("Bundle")
112            .field("path", &self.path)
113            .field("quantization", &self.quantization)
114            .field("tensors", &self.tensors.len())
115            .finish()
116    }
117}
118
119fn read_slice(mmap: &Mmap, start: usize, len: usize) -> Result<&[u8], EngineError> {
120    let end = start
121        .checked_add(len)
122        .ok_or_else(|| EngineError::Format("offset overflow".into()))?;
123    if end > mmap.len() {
124        return Err(EngineError::Format(format!(
125            "offset [{start},{len}] out of range (bin size {})",
126            mmap.len()
127        )));
128    }
129    Ok(&mmap[start..end])
130}
131
132fn offset_pair(v: &Value, key: &str) -> Result<(usize, usize), EngineError> {
133    let arr = v
134        .get(key)
135        .and_then(|x| x.as_array())
136        .ok_or_else(|| EngineError::Format(format!("missing offset {key}")))?;
137    if arr.len() != 2 {
138        return Err(EngineError::Format(format!("bad offset {key}")));
139    }
140    let s = arr[0]
141        .as_u64()
142        .ok_or_else(|| EngineError::Format("offset start".into()))? as usize;
143    let l = arr[1]
144        .as_u64()
145        .ok_or_else(|| EngineError::Format("offset len".into()))? as usize;
146    Ok((s, l))
147}
148
149fn f16_bytes_to_f32(bytes: &[u8]) -> Result<Vec<f32>, EngineError> {
150    if !bytes.len().is_multiple_of(2) {
151        return Err(EngineError::Format("f16 byte length odd".into()));
152    }
153    let mut out = Vec::with_capacity(bytes.len() / 2);
154    for c in bytes.as_chunks::<2>().0 {
155        let h = f16::from_le_bytes(*c);
156        out.push(h.to_f32());
157    }
158    Ok(out)
159}
160
161fn f32_bytes_to_f32(bytes: &[u8]) -> Result<Vec<f32>, EngineError> {
162    if !bytes.len().is_multiple_of(4) {
163        return Err(EngineError::Format("f32 byte length not aligned".into()));
164    }
165    let mut out = Vec::with_capacity(bytes.len() / 4);
166    for c in bytes.as_chunks::<4>().0 {
167        out.push(f32::from_le_bytes(*c));
168    }
169    Ok(out)
170}
171
172pub fn load_bundle(path: impl AsRef<Path>) -> Result<Bundle, EngineError> {
173    let path = path.as_ref();
174    let cfg_path = path.join("config.json");
175    let bin_path = path.join("weight.bin");
176    if !cfg_path.is_file() {
177        return Err(EngineError::Format(format!(
178            "missing config.json in {}",
179            path.display()
180        )));
181    }
182    if !bin_path.is_file() {
183        return Err(EngineError::Format(format!(
184            "missing weight.bin in {}",
185            path.display()
186        )));
187    }
188    let cfg_text = std::fs::read_to_string(&cfg_path)?;
189    let cfg: BundleConfig =
190        serde_json::from_str(&cfg_text).map_err(|e| EngineError::Format(e.to_string()))?;
191    if cfg.format != BUNDLE_FORMAT {
192        return Err(EngineError::Format(format!(
193            "unsupported format {:?}",
194            cfg.format
195        )));
196    }
197    if cfg.format_version != 1 && cfg.format_version != 2 {
198        return Err(EngineError::Format(format!(
199            "unsupported format_version {}",
200            cfg.format_version
201        )));
202    }
203    let file = File::open(&bin_path)?;
204    let mmap = unsafe { Mmap::map(&file)? };
205    let mmap = Arc::new(mmap);
206
207    let mut tensors = HashMap::new();
208    for (name, meta) in &cfg.tensors {
209        let kind = meta
210            .get("kind")
211            .and_then(|v| v.as_str())
212            .ok_or_else(|| EngineError::Format(format!("tensor {name} missing kind")))?;
213        let offsets = meta
214            .get("offsets")
215            .ok_or_else(|| EngineError::Format(format!("tensor {name} missing offsets")))?;
216        match kind {
217            "codebook" => {
218                let bits = meta
219                    .get("bits")
220                    .and_then(|v| v.as_u64())
221                    .ok_or_else(|| EngineError::Quant("bits".into()))?
222                    as u8;
223                if !matches!(bits, 1 | 2 | 3 | 4 | 8) {
224                    return Err(EngineError::Quant(format!("unsupported bits {bits}")));
225                }
226                let group_size =
227                    meta.get("group_size")
228                        .and_then(|v| v.as_u64())
229                        .unwrap_or(cfg.group_size_default as u64) as usize;
230                let shape = meta
231                    .get("shape")
232                    .and_then(|v| v.as_array())
233                    .ok_or_else(|| EngineError::Format("shape".into()))?;
234                if shape.len() != 2 {
235                    return Err(EngineError::Format("codebook shape must be [K,N]".into()));
236                }
237                let k = shape[0].as_u64().unwrap() as usize;
238                let n = shape[1].as_u64().unwrap() as usize;
239                let row_pad = meta.get("row_pad").and_then(|v| v.as_u64()).unwrap_or(0) as usize;
240                let share = meta
241                    .get("codebook_share")
242                    .and_then(|v| v.as_str())
243                    .unwrap_or("group")
244                    .to_string();
245                let (ps, pl) = offset_pair(offsets, "packed_indices")?;
246                let (cs, cl) = offset_pair(offsets, "codebook")?;
247                let packed = read_slice(&mmap, ps, pl)?.to_vec();
248                let cb_raw = read_slice(&mmap, cs, cl)?;
249                let codebook = f16_bytes_to_f32(cb_raw)?;
250                let kc = 1usize << bits;
251                let codebook_shape = if share == "group" {
252                    if !codebook.len().is_multiple_of(kc) {
253                        return Err(EngineError::ShapeMismatch("bad group codebook size".into()));
254                    }
255                    vec![codebook.len() / kc, kc]
256                } else {
257                    if n * kc == 0 || codebook.len() % (n * kc) != 0 {
258                        return Err(EngineError::ShapeMismatch(
259                            "bad channel codebook size".into(),
260                        ));
261                    }
262                    let g = codebook.len() / (n * kc);
263                    vec![g, n, kc]
264                };
265                let hadamard = meta
266                    .get("hadamard")
267                    .cloned()
268                    .unwrap_or(Value::Object(Default::default()));
269                tensors.insert(
270                    name.clone(),
271                    TensorData::Codebook(QuantTensor {
272                        bits,
273                        group_size,
274                        shape: (k, n),
275                        row_pad,
276                        codebook_share: share,
277                        packed_indices: packed,
278                        codebook,
279                        codebook_shape,
280                        hadamard,
281                    }),
282                );
283            }
284            "raw" => {
285                let dtype = meta
286                    .get("dtype")
287                    .and_then(|v| v.as_str())
288                    .unwrap_or("f16")
289                    .to_string();
290                let shape: Vec<usize> = meta
291                    .get("shape")
292                    .and_then(|v| v.as_array())
293                    .ok_or_else(|| EngineError::Format("raw shape".into()))?
294                    .iter()
295                    .map(|x| x.as_u64().unwrap() as usize)
296                    .collect();
297                let (ds, dl) = offset_pair(offsets, "data")?;
298                let raw = read_slice(&mmap, ds, dl)?;
299                let data = if dtype == "f32" {
300                    f32_bytes_to_f32(raw)?
301                } else {
302                    f16_bytes_to_f32(raw)?
303                };
304                tensors.insert(name.clone(), TensorData::Raw { dtype, shape, data });
305            }
306            other => {
307                return Err(EngineError::Format(format!(
308                    "unknown tensor kind {other:?}"
309                )));
310            }
311        }
312    }
313
314    Ok(Bundle {
315        path: path.to_path_buf(),
316        model: cfg.model,
317        quantization: cfg.quantization,
318        group_size_default: cfg.group_size_default,
319        hadamard_seed: cfg.hadamard_seed,
320        tensors,
321        mmap,
322    })
323}
324
325/// Rotated-space reconstruction (matches Python `dequantize`).
326pub fn dequantize(t: &QuantTensor) -> Result<Vec<f32>, EngineError> {
327    let (k0, n) = t.shape;
328    let gs = t.group_size;
329    let kc = 1usize << t.bits;
330    if t.codebook_share == "group" {
331        if t.codebook_shape.len() != 2 {
332            return Err(EngineError::ShapeMismatch(
333                "group codebook must be 2D".into(),
334            ));
335        }
336        let num_groups = t.codebook_shape[0];
337        let k_work = num_groups * gs;
338        let expected = k_work * n;
339        let indices = unpack_indices(&t.packed_indices, expected, t.bits)?;
340        dequant_lookup_group(&indices, &t.codebook, num_groups, gs, n, kc, k0)
341    } else {
342        // channel share
343        if t.codebook_shape.len() != 3 {
344            return Err(EngineError::ShapeMismatch(
345                "channel codebook must be 3D".into(),
346            ));
347        }
348        let num_groups = t.codebook_shape[0];
349        let k_work = num_groups * gs;
350        let expected = k_work * n;
351        let indices = unpack_indices(&t.packed_indices, expected, t.bits)?;
352        let mut out = vec![0.0f32; k_work * n];
353        for g in 0..num_groups {
354            for r in 0..gs {
355                let row = g * gs + r;
356                for j in 0..n {
357                    let idx = indices[row * n + j] as usize;
358                    let base = (g * n + j) * kc;
359                    out[row * n + j] = t.codebook[base + idx];
360                }
361            }
362        }
363        out.truncate(k0 * n);
364        Ok(out)
365    }
366}
367
368impl Bundle {
369    pub fn weight_f32(&self, name: &str) -> Result<Vec<f32>, EngineError> {
370        Ok(self.weight_loaded(name)?.data)
371    }
372
373    /// Load a weight in **original space** (Python `reconstruct_weight`).
374    ///
375    /// Codebook tensors are stored rotated (`W_rot = H@S@W` on axis 0). Fused
376    /// `hdm_linear` unrotates `y = W_rot @ x`, which is valid for dense GEMM but
377    /// **not** for embedding row gather (`e = W[token]`). Unrotating the full
378    /// matrix here makes lookup, tied `lm_head`, and `linear` all match HF inject.
379    pub fn weight_loaded(&self, name: &str) -> Result<LoadedWeight, EngineError> {
380        match self.tensors.get(name) {
381            Some(TensorData::Codebook(q)) => {
382                let t0 = std::time::Instant::now();
383                let mut data = dequantize(q)?;
384                crate::profile::load_profile_add_dequant(crate::profile::elapsed_ms(t0));
385                let applied = q
386                    .hadamard
387                    .get("applied")
388                    .and_then(|v| v.as_bool())
389                    .unwrap_or(false);
390                if applied {
391                    let (k0, n) = q.shape;
392                    if data.len() != k0 * n {
393                        return Err(EngineError::ShapeMismatch(format!(
394                            "dequant {name} len {} != shape {k0}*{n}",
395                            data.len()
396                        )));
397                    }
398                    let seed = self
399                        .hadamard_seed
400                        .or_else(|| q.hadamard.get("seed").and_then(|v| v.as_i64()));
401                    if k0 > 1 {
402                        let t1 = std::time::Instant::now();
403                        let tiles = hadamard_tile_sizes_from_meta(&q.hadamard, k0)?;
404                        hadamard_blocked_rows_tiles(&mut data, k0, n, seed, true, &tiles)?;
405                        crate::profile::load_profile_add_unrotate(crate::profile::elapsed_ms(t1));
406                    }
407                }
408                Ok(LoadedWeight {
409                    data,
410                    hdm_seed: None,
411                })
412            }
413            Some(TensorData::Raw { data, .. }) => Ok(LoadedWeight {
414                data: data.clone(),
415                hdm_seed: None,
416            }),
417            None => Err(EngineError::Format(format!("missing tensor {name}"))),
418        }
419    }
420
421    /// Load the first tensor that exists among `names` (HF vs tiny/blk aliases).
422    pub fn weight_f32_any(&self, names: &[&str]) -> Result<Vec<f32>, EngineError> {
423        Ok(self.weight_loaded_any(names)?.data)
424    }
425
426    pub fn weight_loaded_any(&self, names: &[&str]) -> Result<LoadedWeight, EngineError> {
427        let mut tried = Vec::with_capacity(names.len());
428        for name in names {
429            match self.weight_loaded(name) {
430                Ok(v) => return Ok(v),
431                Err(EngineError::Format(_)) => tried.push(*name),
432                Err(e) => return Err(e),
433            }
434        }
435        Err(EngineError::Format(format!(
436            "missing tensor (tried {})",
437            tried.join(", ")
438        )))
439    }
440}
441
442/// Tile sizes from Python `hadamard.blocks` (`[{start,size},…]`); greedy pow2 if absent/invalid.
443fn hadamard_tile_sizes_from_meta(hadamard: &Value, rows: usize) -> Result<Vec<usize>, EngineError> {
444    if let Some(blocks) = hadamard.get("blocks").and_then(|v| v.as_array()) {
445        if !blocks.is_empty() {
446            let mut sizes = Vec::with_capacity(blocks.len());
447            let mut pos = 0usize;
448            let mut ok = true;
449            for b in blocks {
450                let start = b.get("start").and_then(|x| x.as_u64()).map(|x| x as usize);
451                let size = b.get("size").and_then(|x| x.as_u64()).map(|x| x as usize);
452                match (start, size) {
453                    (Some(s), Some(sz)) if s == pos && sz > 0 => {
454                        sizes.push(sz);
455                        pos = pos.saturating_add(sz);
456                    }
457                    _ => {
458                        ok = false;
459                        break;
460                    }
461                }
462            }
463            if ok && pos == rows {
464                return Ok(sizes);
465            }
466        }
467    }
468    pow2_tile_sizes(rows)
469}
470
471/// Dequantized original-space weight (codebook tensors already unrotated).
472///
473/// `hdm_seed` is kept for API compatibility; `weight_loaded` always clears it
474/// so Session GEMM uses `linear` on reconstructed `W`.
475#[derive(Debug, Clone)]
476pub struct LoadedWeight {
477    pub data: Vec<f32>,
478    pub hdm_seed: Option<i64>,
479}
480
481#[cfg(test)]
482mod tests {
483    use super::*;
484    use crate::fixture::{
485        make_channel_quant_tensor, make_group_quant_tensor, rel_rmse, write_tiny_q4_bundle,
486    };
487    use aria_kernel::hadamard_blocked_rows;
488    use serde_json::json;
489
490    #[test]
491    fn load_and_dequant() {
492        let dir = tempfile::tempdir().unwrap();
493        let (rmse, _) = write_tiny_q4_bundle(dir.path()).unwrap();
494        assert!(rmse < 0.5, "rmse {rmse}");
495        let b = load_bundle(dir.path()).unwrap();
496        assert_eq!(b.model.hidden_size, 64);
497        let w = b.weight_f32("blk.0.attn_q.weight").unwrap();
498        assert_eq!(w.len(), 64 * 64);
499        // Core guarantee: hadamard.applied + codebook K=16 for q4.
500        match b.tensors.get("blk.0.attn_q.weight").unwrap() {
501            TensorData::Codebook(q) => {
502                assert_eq!(q.bits, 4);
503                assert_eq!(1usize << q.bits, q.codebook_shape[1]);
504                assert_eq!(q.hadamard.get("applied"), Some(&json!(true)));
505            }
506            _ => panic!("expected codebook"),
507        }
508        assert!(matches!(
509            b.tensors.get("blk.0.attn_norm.weight"),
510            Some(TensorData::Raw { .. })
511        ));
512    }
513
514    #[test]
515    fn bad_format() {
516        let dir = tempfile::tempdir().unwrap();
517        std::fs::write(dir.path().join("config.json"), r#"{"format":"nope"}"#).unwrap();
518        std::fs::write(dir.path().join("weight.bin"), b"").unwrap();
519        let err = load_bundle(dir.path()).unwrap_err();
520        assert!(matches!(err, EngineError::Format(_)));
521    }
522
523    #[test]
524    fn missing_files() {
525        let dir = tempfile::tempdir().unwrap();
526        assert!(matches!(
527            load_bundle(dir.path()),
528            Err(EngineError::Format(_))
529        ));
530    }
531
532    #[test]
533    fn load_v2_blocked_hadamard_meta() {
534        let dir = tempfile::tempdir().unwrap();
535        write_tiny_q4_bundle(dir.path()).unwrap();
536        let cfg_text = std::fs::read_to_string(dir.path().join("config.json")).unwrap();
537        let cfg: serde_json::Value = serde_json::from_str(&cfg_text).unwrap();
538        assert_eq!(cfg["format_version"], 2);
539        assert_eq!(cfg["hadamard_seed"], 0);
540        let b = load_bundle(dir.path()).unwrap();
541        assert_eq!(b.hadamard_seed, Some(0));
542        match b.tensors.get("blk.0.attn_q.weight").unwrap() {
543            TensorData::Codebook(q) => {
544                assert_eq!(q.hadamard.get("mode"), Some(&json!("blocked")));
545                assert_eq!(q.hadamard.get("applied"), Some(&json!(true)));
546                let blocks = q.hadamard["blocks"].as_array().expect("blocks");
547                assert!(!blocks.is_empty());
548                let k = q.shape.0;
549                let covered: usize = blocks
550                    .iter()
551                    .map(|b| b["size"].as_u64().unwrap() as usize)
552                    .sum();
553                assert_eq!(covered, k);
554                assert_eq!(blocks[0]["start"], 0);
555                // First tile is largest power-of-two ≤ k.
556                let first = blocks[0]["size"].as_u64().unwrap() as usize;
557                assert!(first.is_power_of_two());
558                assert!(first <= k);
559                if k > first {
560                    assert_eq!(blocks[1]["start"], first as u64);
561                }
562            }
563            _ => panic!("expected codebook"),
564        }
565    }
566
567    #[test]
568    fn codebook_weight_loaded_unrotates_like_reconstruct() {
569        let dir = tempfile::tempdir().unwrap();
570        write_tiny_q4_bundle(dir.path()).unwrap();
571        let b = load_bundle(dir.path()).unwrap();
572        let name = "blk.0.attn_q.weight";
573        let q = match b.tensors.get(name).unwrap() {
574            TensorData::Codebook(q) => q,
575            _ => panic!("expected codebook"),
576        };
577        let mut expected = dequantize(q).unwrap();
578        let (k0, n) = q.shape;
579        let seed = b
580            .hadamard_seed
581            .or_else(|| q.hadamard.get("seed").and_then(|v| v.as_i64()));
582        hadamard_blocked_rows(&mut expected, k0, n, seed, true).unwrap();
583        let loaded = b.weight_loaded(name).unwrap();
584        assert!(
585            loaded.hdm_seed.is_none(),
586            "reconstructed weights are original-space; Session uses linear()"
587        );
588        assert_eq!(loaded.data.len(), expected.len());
589        for (a, e) in loaded.data.iter().zip(expected.iter()) {
590            assert!((a - e).abs() < 1e-5, "{a} vs {e}");
591        }
592        // Rotated dequant row must differ from reconstructed row (axis-0 mix).
593        let rotated = dequantize(q).unwrap();
594        let row = n;
595        let rot_norm: f32 = rotated[..row]
596            .iter()
597            .zip(loaded.data[..row].iter())
598            .map(|(a, b)| (a - b) * (a - b))
599            .sum();
600        assert!(
601            rot_norm.sqrt() > 1e-4,
602            "embedding/linear rows must change under blocked unrotate"
603        );
604    }
605
606    #[test]
607    fn load_accepts_format_version_1() {
608        let dir = tempfile::tempdir().unwrap();
609        write_tiny_q4_bundle(dir.path()).unwrap();
610        let cfg_path = dir.path().join("config.json");
611        let mut cfg: serde_json::Value =
612            serde_json::from_str(&std::fs::read_to_string(&cfg_path).unwrap()).unwrap();
613        cfg["format_version"] = json!(1);
614        // Legacy fixtures may omit blocked mode; loader must still accept v1.
615        if let Some(tensors) = cfg["tensors"].as_object_mut() {
616            for meta in tensors.values_mut() {
617                if meta.get("kind") == Some(&json!("codebook")) {
618                    if let Some(h) = meta.get_mut("hadamard") {
619                        if let Some(o) = h.as_object_mut() {
620                            o.remove("mode");
621                            o.remove("blocks");
622                        }
623                    }
624                }
625            }
626        }
627        std::fs::write(&cfg_path, serde_json::to_string_pretty(&cfg).unwrap()).unwrap();
628        let b = load_bundle(dir.path()).unwrap();
629        assert!(!b.tensors.is_empty());
630    }
631
632    #[test]
633    fn load_rejects_format_version_3() {
634        let dir = tempfile::tempdir().unwrap();
635        write_tiny_q4_bundle(dir.path()).unwrap();
636        let cfg_path = dir.path().join("config.json");
637        let mut cfg: serde_json::Value =
638            serde_json::from_str(&std::fs::read_to_string(&cfg_path).unwrap()).unwrap();
639        cfg["format_version"] = json!(3);
640        std::fs::write(&cfg_path, serde_json::to_string_pretty(&cfg).unwrap()).unwrap();
641        let err = load_bundle(dir.path()).unwrap_err();
642        assert!(matches!(err, EngineError::Format(_)));
643        let msg = format!("{err}");
644        assert!(msg.contains("format_version"), "{msg}");
645    }
646
647    /// Mirror model `test_quant.test_dequant_error_bounds` (linspace stand-in; Spec-ish bands).
648    #[test]
649    fn dequant_error_bounds_group() {
650        let mut rng = 0u64;
651        let mut randn = || {
652            rng = rng.wrapping_mul(6364136223846793005).wrapping_add(1);
653            let u = ((rng >> 33) as f32) / (u32::MAX as f32);
654            (u - 0.5) * 2.0
655        };
656        let k = 64usize;
657        let n = 16usize;
658        let mut w = vec![0.0f32; k * n];
659        for v in &mut w {
660            *v = randn();
661        }
662        // Linspace is coarser than Lloyd-Max; keep Spec-aligned upper bands.
663        let bounds = [(8u8, 0.25f32), (4, 0.45), (3, 0.60), (2, 0.85), (1, 1.20)];
664        for (bits, lim) in bounds {
665            let t = make_group_quant_tensor(&w, k, n, 32, bits);
666            assert_eq!(t.codebook_shape[1], 1usize << bits);
667            assert_eq!(t.hadamard.get("applied"), Some(&json!(true)));
668            let recon = dequantize(&t).unwrap();
669            assert_eq!(recon.len(), k * n);
670            let err = rel_rmse(&w, &recon);
671            assert!(err <= lim, "q{bits} group rel_rmse={err} > {lim}");
672        }
673    }
674
675    #[test]
676    fn dequant_channel_q4_tighter() {
677        let mut rng = 1u64;
678        let mut randn = || {
679            rng = rng.wrapping_mul(6364136223846793005).wrapping_add(1);
680            let u = ((rng >> 33) as f32) / (u32::MAX as f32);
681            (u - 0.5) * 2.0
682        };
683        let k = 64usize;
684        let n = 16usize;
685        let mut w = vec![0.0f32; k * n];
686        for v in &mut w {
687            *v = randn();
688        }
689        let t = make_channel_quant_tensor(&w, k, n, 32, 4);
690        assert_eq!(t.codebook_shape, vec![2, n, 16]);
691        let g = make_group_quant_tensor(&w, k, n, 32, 4);
692        assert!(t.codebook.len() > g.codebook.len() * 8);
693        let recon = dequantize(&t).unwrap();
694        let err = rel_rmse(&w, &recon);
695        assert!(err <= 0.35, "q4 channel rel_rmse={err}");
696    }
697
698    #[test]
699    fn dequant_bad_share_shape() {
700        let mut t = make_group_quant_tensor(&[1.0, 2.0, 3.0, 4.0], 2, 2, 2, 4);
701        t.codebook_share = "channel".into(); // shape still 2D → error
702        assert!(matches!(dequantize(&t), Err(EngineError::ShapeMismatch(_))));
703    }
704
705    /// Optional golden path: `ARIA_TINY_BUNDLE` → model `--tiny` export directory.
706    #[test]
707    fn load_aria_tiny_bundle_from_env() {
708        let Ok(path) = std::env::var("ARIA_TINY_BUNDLE") else {
709            return;
710        };
711        let b = load_bundle(&path).expect("ARIA_TINY_BUNDLE must be a valid aria-quant-bundle");
712        assert_eq!(b.quantization.chars().next(), Some('q'));
713        assert!(b.model.hidden_size > 0);
714        assert!(!b.tensors.is_empty());
715        // At least one codebook tensor dequants to declared shape.
716        let (name, q) = b
717            .tensors
718            .iter()
719            .find_map(|(n, t)| match t {
720                TensorData::Codebook(q) => Some((n, q)),
721                _ => None,
722            })
723            .expect("bundle has codebook tensors");
724        let recon = dequantize(q).unwrap();
725        assert_eq!(recon.len(), q.shape.0 * q.shape.1, "{name}");
726        assert_eq!(q.hadamard.get("applied"), Some(&json!(true)), "{name}");
727    }
728
729    #[test]
730    fn hadamard_tiles_prefer_bundle_blocks() {
731        let greedy = pow2_tile_sizes(10).unwrap();
732        assert_eq!(greedy, vec![8, 2]);
733        let meta = json!({
734            "applied": true,
735            "mode": "blocked",
736            "blocks": [{"start": 0, "size": 8}, {"start": 8, "size": 2}]
737        });
738        assert_eq!(hadamard_tile_sizes_from_meta(&meta, 10).unwrap(), greedy);
739        let bad = json!({"blocks": [{"start": 0, "size": 4}]});
740        assert_eq!(hadamard_tile_sizes_from_meta(&bad, 10).unwrap(), greedy);
741        let empty = json!({});
742        assert_eq!(hadamard_tile_sizes_from_meta(&empty, 10).unwrap(), greedy);
743    }
744
745    #[test]
746    fn gemma4_embed_and_ple_codebook_row_gather() {
747        // Mirrors model quantize_weight: axis-0 rotate, group codebook, LSB pack,
748        // shape [vocab, cols] — engine must unrotate then gather token rows.
749        use half::f16;
750        let vocab = 10usize;
751        let hidden = 8usize;
752        let packed_ple = 12usize; // 3 layers * 4
753        let gs = 8usize;
754        let seed = Some(0i64);
755
756        let mut emb: Vec<f32> = (0..vocab * hidden)
757            .map(|i| (i as f32) * 0.01 - 0.05)
758            .collect();
759        let mut ple: Vec<f32> = (0..vocab * packed_ple)
760            .map(|i| (i as f32) * 0.003 - 0.02)
761            .collect();
762        let emb_orig = emb.clone();
763        let ple_orig = ple.clone();
764        hadamard_blocked_rows(&mut emb, vocab, hidden, seed, false).unwrap();
765        hadamard_blocked_rows(&mut ple, vocab, packed_ple, seed, false).unwrap();
766
767        let write_cb = |name: &str,
768                        w_rot: &[f32],
769                        k: usize,
770                        n: usize,
771                        bin: &mut Vec<u8>,
772                        tensors: &mut serde_json::Map<String, Value>| {
773            let t = make_group_quant_tensor(w_rot, k, n, gs, 4);
774            let pi_s = bin.len();
775            bin.extend_from_slice(&t.packed_indices);
776            let pi_l = bin.len() - pi_s;
777            let cb_s = bin.len();
778            for &v in &t.codebook {
779                bin.extend_from_slice(&f16::from_f32(v).to_le_bytes());
780            }
781            let cb_l = bin.len() - cb_s;
782            let mut blocks = Vec::new();
783            let mut start = 0usize;
784            for sz in pow2_tile_sizes(k).unwrap() {
785                blocks.push(json!({"start": start, "size": sz}));
786                start += sz;
787            }
788            tensors.insert(
789                name.to_string(),
790                json!({
791                    "kind": "codebook",
792                    "bits": 4,
793                    "group_size": gs,
794                    "shape": [k, n],
795                    "row_pad": 0,
796                    "codebook_share": "group",
797                    "hadamard": {
798                        "applied": true,
799                        "axis": 0,
800                        "seed": 0,
801                        "mode": "blocked",
802                        "blocks": blocks
803                    },
804                    "offsets": {
805                        "packed_indices": [pi_s, pi_l],
806                        "codebook": [cb_s, cb_l]
807                    }
808                }),
809            );
810        };
811
812        let dir = tempfile::tempdir().unwrap();
813        let mut bin = Vec::new();
814        let mut tensors = serde_json::Map::new();
815        write_cb(
816            "model.language_model.embed_tokens.weight",
817            &emb,
818            vocab,
819            hidden,
820            &mut bin,
821            &mut tensors,
822        );
823        write_cb(
824            "model.language_model.embed_tokens_per_layer.weight",
825            &ple,
826            vocab,
827            packed_ple,
828            &mut bin,
829            &mut tensors,
830        );
831        let cfg = json!({
832            "format": "aria-quant-bundle",
833            "format_version": 2,
834            "quantization": "q4",
835            "group_size_default": gs,
836            "hadamard_seed": 0,
837            "model": {
838                "hidden_size": hidden,
839                "num_layers": 3,
840                "num_attention_heads": 2,
841                "num_kv_heads": 1,
842                "intermediate_size": 16,
843                "vocab_size": vocab,
844                "context_length": 32,
845                "rope_theta": 10000.0
846            },
847            "tensors": tensors
848        });
849        std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
850        std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
851
852        let b = load_bundle(dir.path()).unwrap();
853        let loaded_emb = b
854            .weight_loaded("model.language_model.embed_tokens.weight")
855            .unwrap();
856        let loaded_ple = b
857            .weight_loaded("model.language_model.embed_tokens_per_layer.weight")
858            .unwrap();
859        assert_eq!(loaded_emb.data.len(), vocab * hidden);
860        assert_eq!(loaded_ple.data.len(), vocab * packed_ple);
861        // Quant is lossy; unrotate+gather must match reconstruct (dequant+unrotate).
862        for (name, cols, orig, loaded) in [
863            (
864                "embed",
865                hidden,
866                emb_orig.as_slice(),
867                loaded_emb.data.as_slice(),
868            ),
869            (
870                "ple",
871                packed_ple,
872                ple_orig.as_slice(),
873                loaded_ple.data.as_slice(),
874            ),
875        ] {
876            let q = match b.tensors.get(match name {
877                "embed" => "model.language_model.embed_tokens.weight",
878                _ => "model.language_model.embed_tokens_per_layer.weight",
879            }) {
880                Some(TensorData::Codebook(q)) => q,
881                _ => panic!("{name}"),
882            };
883            let mut recon = dequantize(q).unwrap();
884            hadamard_blocked_rows(&mut recon, vocab, cols, seed, true).unwrap();
885            for (a, e) in loaded.iter().zip(recon.iter()) {
886                assert!((a - e).abs() < 1e-5, "{name} {a} vs {e}");
887            }
888            for tid in [0usize, 2, 9] {
889                let row_l = &loaded[tid * cols..(tid + 1) * cols];
890                let row_r = &recon[tid * cols..(tid + 1) * cols];
891                for (a, e) in row_l.iter().zip(row_r.iter()) {
892                    assert!((a - e).abs() < 1e-5, "{name} tid={tid}");
893                }
894                let orig_row = &orig[tid * cols..(tid + 1) * cols];
895                assert!(
896                    orig_row.iter().any(|v| v.abs() > 1e-6),
897                    "{name} orig row {tid} unexpectedly zero"
898                );
899            }
900        }
901    }
902}