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 #[serde(default)]
45 pub conv_l_cache: Option<usize>,
46 #[serde(default)]
48 pub sliding_window: Option<usize>,
49 #[serde(default)]
51 pub partial_rotary_factor: Option<f32>,
52 #[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
325pub 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 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 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 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
442fn 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#[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 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 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 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 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 #[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 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(); assert!(matches!(dequantize(&t), Err(EngineError::ShapeMismatch(_))));
703 }
704
705 #[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 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 use half::f16;
750 let vocab = 10usize;
751 let hidden = 8usize;
752 let packed_ple = 12usize; 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 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}