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(
444    hadamard: &Value,
445    rows: usize,
446) -> Result<Vec<usize>, EngineError> {
447    if let Some(blocks) = hadamard.get("blocks").and_then(|v| v.as_array()) {
448        if !blocks.is_empty() {
449            let mut sizes = Vec::with_capacity(blocks.len());
450            let mut pos = 0usize;
451            let mut ok = true;
452            for b in blocks {
453                let start = b.get("start").and_then(|x| x.as_u64()).map(|x| x as usize);
454                let size = b.get("size").and_then(|x| x.as_u64()).map(|x| x as usize);
455                match (start, size) {
456                    (Some(s), Some(sz)) if s == pos && sz > 0 => {
457                        sizes.push(sz);
458                        pos = pos.saturating_add(sz);
459                    }
460                    _ => {
461                        ok = false;
462                        break;
463                    }
464                }
465            }
466            if ok && pos == rows {
467                return Ok(sizes);
468            }
469        }
470    }
471    pow2_tile_sizes(rows)
472}
473
474/// Dequantized original-space weight (codebook tensors already unrotated).
475///
476/// `hdm_seed` is kept for API compatibility; `weight_loaded` always clears it
477/// so Session GEMM uses `linear` on reconstructed `W`.
478#[derive(Debug, Clone)]
479pub struct LoadedWeight {
480    pub data: Vec<f32>,
481    pub hdm_seed: Option<i64>,
482}
483
484#[cfg(test)]
485mod tests {
486    use super::*;
487    use crate::fixture::{
488        make_channel_quant_tensor, make_group_quant_tensor, rel_rmse, write_tiny_q4_bundle,
489    };
490    use aria_kernel::hadamard_blocked_rows;
491    use serde_json::json;
492
493    #[test]
494    fn load_and_dequant() {
495        let dir = tempfile::tempdir().unwrap();
496        let (rmse, _) = write_tiny_q4_bundle(dir.path()).unwrap();
497        assert!(rmse < 0.5, "rmse {rmse}");
498        let b = load_bundle(dir.path()).unwrap();
499        assert_eq!(b.model.hidden_size, 64);
500        let w = b.weight_f32("blk.0.attn_q.weight").unwrap();
501        assert_eq!(w.len(), 64 * 64);
502        // Core guarantee: hadamard.applied + codebook K=16 for q4.
503        match b.tensors.get("blk.0.attn_q.weight").unwrap() {
504            TensorData::Codebook(q) => {
505                assert_eq!(q.bits, 4);
506                assert_eq!(1usize << q.bits, q.codebook_shape[1]);
507                assert_eq!(q.hadamard.get("applied"), Some(&json!(true)));
508            }
509            _ => panic!("expected codebook"),
510        }
511        assert!(matches!(
512            b.tensors.get("blk.0.attn_norm.weight"),
513            Some(TensorData::Raw { .. })
514        ));
515    }
516
517    #[test]
518    fn bad_format() {
519        let dir = tempfile::tempdir().unwrap();
520        std::fs::write(dir.path().join("config.json"), r#"{"format":"nope"}"#).unwrap();
521        std::fs::write(dir.path().join("weight.bin"), b"").unwrap();
522        let err = load_bundle(dir.path()).unwrap_err();
523        assert!(matches!(err, EngineError::Format(_)));
524    }
525
526    #[test]
527    fn missing_files() {
528        let dir = tempfile::tempdir().unwrap();
529        assert!(matches!(
530            load_bundle(dir.path()),
531            Err(EngineError::Format(_))
532        ));
533    }
534
535    #[test]
536    fn load_v2_blocked_hadamard_meta() {
537        let dir = tempfile::tempdir().unwrap();
538        write_tiny_q4_bundle(dir.path()).unwrap();
539        let cfg_text = std::fs::read_to_string(dir.path().join("config.json")).unwrap();
540        let cfg: serde_json::Value = serde_json::from_str(&cfg_text).unwrap();
541        assert_eq!(cfg["format_version"], 2);
542        assert_eq!(cfg["hadamard_seed"], 0);
543        let b = load_bundle(dir.path()).unwrap();
544        assert_eq!(b.hadamard_seed, Some(0));
545        match b.tensors.get("blk.0.attn_q.weight").unwrap() {
546            TensorData::Codebook(q) => {
547                assert_eq!(q.hadamard.get("mode"), Some(&json!("blocked")));
548                assert_eq!(q.hadamard.get("applied"), Some(&json!(true)));
549                let blocks = q.hadamard["blocks"].as_array().expect("blocks");
550                assert!(!blocks.is_empty());
551                let k = q.shape.0;
552                let covered: usize = blocks
553                    .iter()
554                    .map(|b| b["size"].as_u64().unwrap() as usize)
555                    .sum();
556                assert_eq!(covered, k);
557                assert_eq!(blocks[0]["start"], 0);
558                // First tile is largest power-of-two ≤ k.
559                let first = blocks[0]["size"].as_u64().unwrap() as usize;
560                assert!(first.is_power_of_two());
561                assert!(first <= k);
562                if k > first {
563                    assert_eq!(blocks[1]["start"], first as u64);
564                }
565            }
566            _ => panic!("expected codebook"),
567        }
568    }
569
570    #[test]
571    fn codebook_weight_loaded_unrotates_like_reconstruct() {
572        let dir = tempfile::tempdir().unwrap();
573        write_tiny_q4_bundle(dir.path()).unwrap();
574        let b = load_bundle(dir.path()).unwrap();
575        let name = "blk.0.attn_q.weight";
576        let q = match b.tensors.get(name).unwrap() {
577            TensorData::Codebook(q) => q,
578            _ => panic!("expected codebook"),
579        };
580        let mut expected = dequantize(q).unwrap();
581        let (k0, n) = q.shape;
582        let seed = b
583            .hadamard_seed
584            .or_else(|| q.hadamard.get("seed").and_then(|v| v.as_i64()));
585        hadamard_blocked_rows(&mut expected, k0, n, seed, true).unwrap();
586        let loaded = b.weight_loaded(name).unwrap();
587        assert!(
588            loaded.hdm_seed.is_none(),
589            "reconstructed weights are original-space; Session uses linear()"
590        );
591        assert_eq!(loaded.data.len(), expected.len());
592        for (a, e) in loaded.data.iter().zip(expected.iter()) {
593            assert!((a - e).abs() < 1e-5, "{a} vs {e}");
594        }
595        // Rotated dequant row must differ from reconstructed row (axis-0 mix).
596        let rotated = dequantize(q).unwrap();
597        let row = n;
598        let rot_norm: f32 = rotated[..row]
599            .iter()
600            .zip(loaded.data[..row].iter())
601            .map(|(a, b)| (a - b) * (a - b))
602            .sum();
603        assert!(
604            rot_norm.sqrt() > 1e-4,
605            "embedding/linear rows must change under blocked unrotate"
606        );
607    }
608
609    #[test]
610    fn load_accepts_format_version_1() {
611        let dir = tempfile::tempdir().unwrap();
612        write_tiny_q4_bundle(dir.path()).unwrap();
613        let cfg_path = dir.path().join("config.json");
614        let mut cfg: serde_json::Value =
615            serde_json::from_str(&std::fs::read_to_string(&cfg_path).unwrap()).unwrap();
616        cfg["format_version"] = json!(1);
617        // Legacy fixtures may omit blocked mode; loader must still accept v1.
618        if let Some(tensors) = cfg["tensors"].as_object_mut() {
619            for meta in tensors.values_mut() {
620                if meta.get("kind") == Some(&json!("codebook")) {
621                    if let Some(h) = meta.get_mut("hadamard") {
622                        if let Some(o) = h.as_object_mut() {
623                            o.remove("mode");
624                            o.remove("blocks");
625                        }
626                    }
627                }
628            }
629        }
630        std::fs::write(&cfg_path, serde_json::to_string_pretty(&cfg).unwrap()).unwrap();
631        let b = load_bundle(dir.path()).unwrap();
632        assert!(!b.tensors.is_empty());
633    }
634
635    #[test]
636    fn load_rejects_format_version_3() {
637        let dir = tempfile::tempdir().unwrap();
638        write_tiny_q4_bundle(dir.path()).unwrap();
639        let cfg_path = dir.path().join("config.json");
640        let mut cfg: serde_json::Value =
641            serde_json::from_str(&std::fs::read_to_string(&cfg_path).unwrap()).unwrap();
642        cfg["format_version"] = json!(3);
643        std::fs::write(&cfg_path, serde_json::to_string_pretty(&cfg).unwrap()).unwrap();
644        let err = load_bundle(dir.path()).unwrap_err();
645        assert!(matches!(err, EngineError::Format(_)));
646        let msg = format!("{err}");
647        assert!(msg.contains("format_version"), "{msg}");
648    }
649
650    /// Mirror model `test_quant.test_dequant_error_bounds` (linspace stand-in; Spec-ish bands).
651    #[test]
652    fn dequant_error_bounds_group() {
653        let mut rng = 0u64;
654        let mut randn = || {
655            rng = rng.wrapping_mul(6364136223846793005).wrapping_add(1);
656            let u = ((rng >> 33) as f32) / (u32::MAX as f32);
657            (u - 0.5) * 2.0
658        };
659        let k = 64usize;
660        let n = 16usize;
661        let mut w = vec![0.0f32; k * n];
662        for v in &mut w {
663            *v = randn();
664        }
665        // Linspace is coarser than Lloyd-Max; keep Spec-aligned upper bands.
666        let bounds = [(8u8, 0.25f32), (4, 0.45), (3, 0.60), (2, 0.85), (1, 1.20)];
667        for (bits, lim) in bounds {
668            let t = make_group_quant_tensor(&w, k, n, 32, bits);
669            assert_eq!(t.codebook_shape[1], 1usize << bits);
670            assert_eq!(t.hadamard.get("applied"), Some(&json!(true)));
671            let recon = dequantize(&t).unwrap();
672            assert_eq!(recon.len(), k * n);
673            let err = rel_rmse(&w, &recon);
674            assert!(err <= lim, "q{bits} group rel_rmse={err} > {lim}");
675        }
676    }
677
678    #[test]
679    fn dequant_channel_q4_tighter() {
680        let mut rng = 1u64;
681        let mut randn = || {
682            rng = rng.wrapping_mul(6364136223846793005).wrapping_add(1);
683            let u = ((rng >> 33) as f32) / (u32::MAX as f32);
684            (u - 0.5) * 2.0
685        };
686        let k = 64usize;
687        let n = 16usize;
688        let mut w = vec![0.0f32; k * n];
689        for v in &mut w {
690            *v = randn();
691        }
692        let t = make_channel_quant_tensor(&w, k, n, 32, 4);
693        assert_eq!(t.codebook_shape, vec![2, n, 16]);
694        let g = make_group_quant_tensor(&w, k, n, 32, 4);
695        assert!(t.codebook.len() > g.codebook.len() * 8);
696        let recon = dequantize(&t).unwrap();
697        let err = rel_rmse(&w, &recon);
698        assert!(err <= 0.35, "q4 channel rel_rmse={err}");
699    }
700
701    #[test]
702    fn dequant_bad_share_shape() {
703        let mut t = make_group_quant_tensor(&[1.0, 2.0, 3.0, 4.0], 2, 2, 2, 4);
704        t.codebook_share = "channel".into(); // shape still 2D → error
705        assert!(matches!(dequantize(&t), Err(EngineError::ShapeMismatch(_))));
706    }
707
708    /// Optional golden path: `ARIA_TINY_BUNDLE` → model `--tiny` export directory.
709    #[test]
710    fn load_aria_tiny_bundle_from_env() {
711        let Ok(path) = std::env::var("ARIA_TINY_BUNDLE") else {
712            return;
713        };
714        let b = load_bundle(&path).expect("ARIA_TINY_BUNDLE must be a valid aria-quant-bundle");
715        assert_eq!(b.quantization.chars().next(), Some('q'));
716        assert!(b.model.hidden_size > 0);
717        assert!(!b.tensors.is_empty());
718        // At least one codebook tensor dequants to declared shape.
719        let (name, q) = b
720            .tensors
721            .iter()
722            .find_map(|(n, t)| match t {
723                TensorData::Codebook(q) => Some((n, q)),
724                _ => None,
725            })
726            .expect("bundle has codebook tensors");
727        let recon = dequantize(q).unwrap();
728        assert_eq!(recon.len(), q.shape.0 * q.shape.1, "{name}");
729        assert_eq!(q.hadamard.get("applied"), Some(&json!(true)), "{name}");
730    }
731
732    #[test]
733    fn hadamard_tiles_prefer_bundle_blocks() {
734        let greedy = pow2_tile_sizes(10).unwrap();
735        assert_eq!(greedy, vec![8, 2]);
736        let meta = json!({
737            "applied": true,
738            "mode": "blocked",
739            "blocks": [{"start": 0, "size": 8}, {"start": 8, "size": 2}]
740        });
741        assert_eq!(hadamard_tile_sizes_from_meta(&meta, 10).unwrap(), greedy);
742        let bad = json!({"blocks": [{"start": 0, "size": 4}]});
743        assert_eq!(hadamard_tile_sizes_from_meta(&bad, 10).unwrap(), greedy);
744        let empty = json!({});
745        assert_eq!(hadamard_tile_sizes_from_meta(&empty, 10).unwrap(), greedy);
746    }
747
748    #[test]
749    fn gemma4_embed_and_ple_codebook_row_gather() {
750        // Mirrors model quantize_weight: axis-0 rotate, group codebook, LSB pack,
751        // shape [vocab, cols] — engine must unrotate then gather token rows.
752        use half::f16;
753        let vocab = 10usize;
754        let hidden = 8usize;
755        let packed_ple = 12usize; // 3 layers * 4
756        let gs = 8usize;
757        let seed = Some(0i64);
758
759        let mut emb: Vec<f32> = (0..vocab * hidden)
760            .map(|i| (i as f32) * 0.01 - 0.05)
761            .collect();
762        let mut ple: Vec<f32> = (0..vocab * packed_ple)
763            .map(|i| (i as f32) * 0.003 - 0.02)
764            .collect();
765        let emb_orig = emb.clone();
766        let ple_orig = ple.clone();
767        hadamard_blocked_rows(&mut emb, vocab, hidden, seed, false).unwrap();
768        hadamard_blocked_rows(&mut ple, vocab, packed_ple, seed, false).unwrap();
769
770        let write_cb = |name: &str,
771                        w_rot: &[f32],
772                        k: usize,
773                        n: usize,
774                        bin: &mut Vec<u8>,
775                        tensors: &mut serde_json::Map<String, Value>| {
776            let t = make_group_quant_tensor(w_rot, k, n, gs, 4);
777            let pi_s = bin.len();
778            bin.extend_from_slice(&t.packed_indices);
779            let pi_l = bin.len() - pi_s;
780            let cb_s = bin.len();
781            for &v in &t.codebook {
782                bin.extend_from_slice(&f16::from_f32(v).to_le_bytes());
783            }
784            let cb_l = bin.len() - cb_s;
785            let mut blocks = Vec::new();
786            let mut start = 0usize;
787            for sz in pow2_tile_sizes(k).unwrap() {
788                blocks.push(json!({"start": start, "size": sz}));
789                start += sz;
790            }
791            tensors.insert(
792                name.to_string(),
793                json!({
794                    "kind": "codebook",
795                    "bits": 4,
796                    "group_size": gs,
797                    "shape": [k, n],
798                    "row_pad": 0,
799                    "codebook_share": "group",
800                    "hadamard": {
801                        "applied": true,
802                        "axis": 0,
803                        "seed": 0,
804                        "mode": "blocked",
805                        "blocks": blocks
806                    },
807                    "offsets": {
808                        "packed_indices": [pi_s, pi_l],
809                        "codebook": [cb_s, cb_l]
810                    }
811                }),
812            );
813        };
814
815        let dir = tempfile::tempdir().unwrap();
816        let mut bin = Vec::new();
817        let mut tensors = serde_json::Map::new();
818        write_cb(
819            "model.language_model.embed_tokens.weight",
820            &emb,
821            vocab,
822            hidden,
823            &mut bin,
824            &mut tensors,
825        );
826        write_cb(
827            "model.language_model.embed_tokens_per_layer.weight",
828            &ple,
829            vocab,
830            packed_ple,
831            &mut bin,
832            &mut tensors,
833        );
834        let cfg = json!({
835            "format": "aria-quant-bundle",
836            "format_version": 2,
837            "quantization": "q4",
838            "group_size_default": gs,
839            "hadamard_seed": 0,
840            "model": {
841                "hidden_size": hidden,
842                "num_layers": 3,
843                "num_attention_heads": 2,
844                "num_kv_heads": 1,
845                "intermediate_size": 16,
846                "vocab_size": vocab,
847                "context_length": 32,
848                "rope_theta": 10000.0
849            },
850            "tensors": tensors
851        });
852        std::fs::write(dir.path().join("config.json"), cfg.to_string()).unwrap();
853        std::fs::write(dir.path().join("weight.bin"), &bin).unwrap();
854
855        let b = load_bundle(dir.path()).unwrap();
856        let loaded_emb = b
857            .weight_loaded("model.language_model.embed_tokens.weight")
858            .unwrap();
859        let loaded_ple = b
860            .weight_loaded("model.language_model.embed_tokens_per_layer.weight")
861            .unwrap();
862        assert_eq!(loaded_emb.data.len(), vocab * hidden);
863        assert_eq!(loaded_ple.data.len(), vocab * packed_ple);
864        // Quant is lossy; unrotate+gather must match reconstruct (dequant+unrotate).
865        for (name, cols, orig, loaded) in [
866            (
867                "embed",
868                hidden,
869                emb_orig.as_slice(),
870                loaded_emb.data.as_slice(),
871            ),
872            (
873                "ple",
874                packed_ple,
875                ple_orig.as_slice(),
876                loaded_ple.data.as_slice(),
877            ),
878        ] {
879            let q = match b.tensors.get(match name {
880                "embed" => "model.language_model.embed_tokens.weight",
881                _ => "model.language_model.embed_tokens_per_layer.weight",
882            }) {
883                Some(TensorData::Codebook(q)) => q,
884                _ => panic!("{name}"),
885            };
886            let mut recon = dequantize(q).unwrap();
887            hadamard_blocked_rows(&mut recon, vocab, cols, seed, true).unwrap();
888            for (a, e) in loaded.iter().zip(recon.iter()) {
889                assert!((a - e).abs() < 1e-5, "{name} {a} vs {e}");
890            }
891            for tid in [0usize, 2, 9] {
892                let row_l = &loaded[tid * cols..(tid + 1) * cols];
893                let row_r = &recon[tid * cols..(tid + 1) * cols];
894                for (a, e) in row_l.iter().zip(row_r.iter()) {
895                    assert!((a - e).abs() < 1e-5, "{name} tid={tid}");
896                }
897                let orig_row = &orig[tid * cols..(tid + 1) * cols];
898                assert!(
899                    orig_row.iter().any(|v| v.abs() > 1e-6),
900                    "{name} orig row {tid} unexpectedly zero"
901                );
902            }
903        }
904    }
905}