1use burn::tensor::{backend::Backend, Tensor};
10
11pub fn fast_walsh_hadamard<B: Backend>(x: Tensor<B, 2>) -> Tensor<B, 2> {
12 let [n, d] = x.dims();
13 let p = d.next_power_of_two();
14 let dev = x.device();
15 let (x_pad, was_padded) = if p != d {
16 (
17 Tensor::cat(vec![x, Tensor::zeros([n, p - d], &dev)], 1),
18 true,
19 )
20 } else {
21 (x, false)
22 };
23 let mut h = 1usize;
24 let mut out = x_pad;
25 while h < p {
26 let step = 2 * h;
27 let r = out.reshape([n, p / step, 2, h]);
28 let l = r
29 .clone()
30 .slice([0..n, 0..(p / step), 0..1, 0..h])
31 .squeeze_dim::<3>(2);
32 let rt = r
33 .slice([0..n, 0..(p / step), 1..2, 0..h])
34 .squeeze_dim::<3>(2);
35 out = Tensor::cat(
36 vec![
37 (l.clone() + rt.clone()).unsqueeze_dim::<4>(2),
38 (l - rt).unsqueeze_dim::<4>(2),
39 ],
40 2,
41 )
42 .reshape([n, p]);
43 h = step;
44 }
45 let out = out.div_scalar((p as f32).sqrt());
46 if was_padded {
47 out.slice([0..n, 0..d])
48 } else {
49 out
50 }
51}
52
53pub fn weight_quant_ternary<B: Backend>(w: Tensor<B, 2>) -> Tensor<B, 2> {
54 let scale = w.clone().abs().mean().unsqueeze_dims(&[0, 0]);
55 let mean = w.clone().mean().unsqueeze_dims(&[0, 0]);
56 let u = w.clone().sub(mean).sign().mul(scale);
57 let base = w.clone().detach();
58 w.add(u.sub(base))
59}
60
61pub fn activation_quant_8bit<B: Backend>(x: Tensor<B, 3>) -> Tensor<B, 3> {
62 let [b, t, d] = x.dims();
63 let flat = x.clone().reshape([b * t, d]);
64 let scale = flat.clone().abs().max_dim(1).clamp_min(1e-5);
65 let y = flat
66 .div(scale.clone())
67 .mul_scalar(127.0)
68 .round()
69 .clamp(-128.0, 127.0)
70 .div_scalar(127.0)
71 .mul(scale)
72 .reshape([b, t, d]);
73 let base = x.clone().detach();
74 x.add(y.sub(base))
75}
76
77pub fn bitnet_v2_quantize<B: Backend>(x: Tensor<B, 3>, bits: usize) -> Tensor<B, 3> {
78 if bits >= 16 {
79 return x;
80 }
81 let [b, t, d] = x.dims();
82 let flat = x.clone().reshape([b * t, d]);
83 let rotated = fast_walsh_hadamard(flat);
84 let deq_rot = if bits >= 8 {
85 let scale = rotated.clone().abs().max_dim(1).clamp_min(1e-12);
86 let q = rotated
87 .clone()
88 .div(scale.clone())
89 .mul_scalar(127.0)
90 .round()
91 .clamp(-128.0, 127.0);
92 fast_walsh_hadamard(q.div(scale).mul_scalar(1.0 / 127.0))
93 } else {
94 let scale = rotated.clone().abs().mean_dim(1).clamp_min(1e-12);
95 let q = rotated
96 .clone()
97 .div(scale.clone())
98 .mul_scalar(7.0)
99 .round()
100 .clamp(-8.0, 7.0);
101 fast_walsh_hadamard(q.div(scale).mul_scalar(1.0 / 7.0))
102 };
103 let y = deq_rot.reshape([b, t, d]);
104 let base = x.clone().detach();
105 x.add(y.sub(base))
106}
107
108pub fn quantize_4bit<B: Backend>(x: Tensor<B, 3>) -> Tensor<B, 3> {
109 bitnet_v2_quantize(x, 4)
110}
111pub fn quantize_8bit<B: Backend>(x: Tensor<B, 3>) -> Tensor<B, 3> {
112 bitnet_v2_quantize(x, 8)
113}
114
115#[cfg(test)]
116mod tests {
117 use super::*;
118 use burn::tensor::Distribution;
119 use burn_ndarray::{NdArray, NdArrayDevice};
120 type B = NdArray;
121 fn dev() -> NdArrayDevice {
122 NdArrayDevice::default()
123 }
124
125 #[test]
126 fn hadamard_roundtrip() {
127 let x = Tensor::<B, 2>::ones([4, 8], &dev());
128 let h2x = fast_walsh_hadamard(fast_walsh_hadamard(x));
129 let v: Vec<f32> = h2x
130 .into_data()
131 .bytes
132 .chunks_exact(4)
133 .map(|b| f32::from_le_bytes(b.try_into().unwrap()))
134 .collect();
135 for (i, val) in v.iter().enumerate() {
136 assert!((val - 1.0).abs() < 0.1, "idx {i}: {val}");
137 }
138 }
139
140 #[test]
141 fn hadamard_non_power_of_two() {
142 let h2x =
143 fast_walsh_hadamard::<B>(fast_walsh_hadamard(Tensor::<B, 2>::ones([2, 7], &dev())));
144 assert_eq!(h2x.dims(), [2, 7]);
145 }
146
147 #[test]
148 fn weight_ternary_values() {
149 let q = weight_quant_ternary(Tensor::<B, 2>::random(
150 [4, 16],
151 Distribution::Normal(0.0, 0.5),
152 &dev(),
153 ));
154 let v: Vec<f32> = q
155 .into_data()
156 .bytes
157 .chunks_exact(4)
158 .map(|b| f32::from_le_bytes(b.try_into().unwrap()))
159 .collect();
160 let s = v.iter().map(|x| x.abs()).sum::<f32>() / v.len() as f32;
161 for val in v {
162 assert!(
163 (val.abs() - s).abs() < 0.01 || val.abs() < 0.01,
164 "{val} not ternary"
165 );
166 }
167 }
168
169 #[test]
170 fn activation_8bit_shape() {
171 assert_eq!(
172 activation_quant_8bit(Tensor::<B, 3>::random(
173 [2, 8, 64],
174 Distribution::Normal(0.0, 1.0),
175 &dev()
176 ))
177 .dims(),
178 [2, 8, 64]
179 );
180 }
181
182 #[test]
183 fn v2_4bit_finite() {
184 let q = quantize_4bit(Tensor::<B, 3>::random(
185 [1, 8, 32],
186 Distribution::Normal(0.0, 1.0),
187 &dev(),
188 ));
189 assert!(q
190 .into_data()
191 .bytes
192 .chunks_exact(4)
193 .map(|b| f32::from_le_bytes(b.try_into().unwrap()))
194 .all(|v| v.is_finite()));
195 }
196
197 #[test]
198 fn v2_8bit_finite() {
199 let q = quantize_8bit(Tensor::<B, 3>::random(
200 [2, 16, 128],
201 Distribution::Normal(0.0, 1.0),
202 &dev(),
203 ));
204 assert!(q
205 .into_data()
206 .bytes
207 .chunks_exact(4)
208 .map(|b| f32::from_le_bytes(b.try_into().unwrap()))
209 .all(|v| v.is_finite()));
210 }
211
212 #[test]
213 fn v2_pass_through() {
214 let x = Tensor::<B, 3>::random([1, 4, 16], Distribution::Normal(0.0, 1.0), &dev());
215 let q = bitnet_v2_quantize(x.clone(), 16);
216 let d: Vec<f32> = (x - q)
217 .into_data()
218 .bytes
219 .chunks_exact(4)
220 .map(|b| f32::from_le_bytes(b.try_into().unwrap()))
221 .collect();
222 assert!(d.iter().all(|&v| v.abs() < 1e-5));
223 }
224}