Skip to main content

burn_bitnet/
lib.rs

1//! # burn-bitnet — BitNet quantization family for Burn
2//!
3//! | Function | Reference | What |
4//! |----------|-----------|------|
5//! | `weight_quant_ternary` | b1.58 (2024) | W -> +/-scale, mean-based |
6//! | `activation_quant_8bit` | b1.58 (2024) | absmax -> 8-bit |
7//! | `bitnet_v2_quantize` | v2 (2025) | Hadamard + absmax/absmean |
8//! | `fast_walsh_hadamard` | v2 (2025) | FWHT O(n log n) |
9use 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}