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(
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#[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 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 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 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 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 #[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 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(); assert!(matches!(dequantize(&t), Err(EngineError::ShapeMismatch(_))));
706 }
707
708 #[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 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 use half::f16;
753 let vocab = 10usize;
754 let hidden = 8usize;
755 let packed_ple = 12usize; 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 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}