1use burn::tensor::{Int, Tensor, backend::Backend};
20
21pub const DEFAULT_Q4_GROUP_SIZE: usize = 32;
23
24pub fn dequantize_q4<B: Backend>(
37 packed: Tensor<B, 2, Int>,
38 scales: Tensor<B, 2>,
39 group_size: usize,
40) -> Tensor<B, 2> {
41 let [rows, packed_cols] = packed.dims();
42 let half = group_size / 2;
43 assert_eq!(
44 packed_cols % half,
45 0,
46 "packed width {packed_cols} not a multiple of half group {half}"
47 );
48 let groups_per_row = packed_cols / half;
49 let cols = groups_per_row * group_size;
50 let [s_rows, s_cols] = scales.dims();
51 assert_eq!(
52 (s_rows, s_cols),
53 (rows, groups_per_row),
54 "scales shape [{s_rows}, {s_cols}] does not match [rows, cols/group] = [{rows}, {groups_per_row}]"
55 );
56
57 let nibbles = packed.reshape([rows, groups_per_row, half]);
59 let lo = nibbles.clone().remainder_scalar(16); let hi = nibbles.div_scalar(16); let nibbles = Tensor::cat(vec![lo, hi], 2).float();
64
65 let scales = scales
67 .unsqueeze_dim::<3>(2)
68 .expand([rows, groups_per_row, group_size]);
69 let w = (nibbles - 8.0) * scales;
70 w.reshape([rows, cols])
71}
72
73#[cfg(test)]
74mod tests {
75 use super::*;
76 use burn::tensor::TensorData;
77
78 type B = burn::backend::NdArray<f32>;
79
80 fn reference_block(bytes: &[u8; 16], scale: f32) -> [f32; 32] {
82 let mut out = [0.0f32; 32];
83 for j in 0..16 {
84 out[j] = ((bytes[j] & 0x0F) as f32 - 8.0) * scale;
85 out[j + 16] = ((bytes[j] >> 4) as f32 - 8.0) * scale;
86 }
87 out
88 }
89
90 #[test]
91 fn dequantize_matches_scalar_reference() {
92 let device = Default::default();
93 let packed_bytes: Vec<i32> = (0..64).map(|i| ((i * 37 + 11) % 256) as i32).collect();
95 let scales: Vec<f32> = vec![0.5, -1.25, 2.0, 0.75];
96 let packed = Tensor::<B, 2, Int>::from_data(
97 TensorData::new(packed_bytes.clone(), [2, 32]),
98 &device,
99 );
100 let scales_t = Tensor::<B, 2>::from_data(TensorData::new(scales.clone(), [2, 2]), &device);
101
102 let w = dequantize_q4(packed, scales_t, DEFAULT_Q4_GROUP_SIZE);
103 let got: Vec<f32> = w.into_data().to_vec().unwrap();
104
105 let mut expected = Vec::with_capacity(128);
106 for row in 0..2 {
107 for g in 0..2 {
108 let mut block = [0u8; 16];
109 for j in 0..16 {
110 block[j] = packed_bytes[row * 32 + g * 16 + j] as u8;
111 }
112 expected.extend_from_slice(&reference_block(&block, scales[row * 2 + g]));
113 }
114 }
115 assert_eq!(got.len(), expected.len());
116 for (g, e) in got.iter().zip(expected.iter()) {
117 assert!((g - e).abs() < 1e-6, "got {g}, expected {e}");
118 }
119 }
120
121 #[test]
122 fn all_eights_pack_to_zero_weights() {
123 let device = Default::default();
124 let packed = Tensor::<B, 2, Int>::from_data(
126 TensorData::new(vec![0x88i32; 16], [1, 16]),
127 &device,
128 );
129 let scales = Tensor::<B, 2>::from_data(TensorData::new(vec![3.5f32], [1, 1]), &device);
130 let w = dequantize_q4(packed, scales, DEFAULT_Q4_GROUP_SIZE);
131 let got: Vec<f32> = w.into_data().to_vec().unwrap();
132 assert_eq!(got, vec![0.0; 32]);
133 }
134}