burn-bitnet 0.1.0

BitNet quantization family for Burn — b1.58 ternary weights + v2 Hadamard activations
Documentation
//! # burn-bitnet — BitNet quantization family for Burn
//!
//! | Function | Reference | What |
//! |----------|-----------|------|
//! | `weight_quant_ternary` | b1.58 (2024) | W -> +/-scale, mean-based |
//! | `activation_quant_8bit` | b1.58 (2024) | absmax -> 8-bit |
//! | `bitnet_v2_quantize` | v2 (2025) | Hadamard + absmax/absmean |
//! | `fast_walsh_hadamard` | v2 (2025) | FWHT O(n log n) |
use burn::tensor::{backend::Backend, Tensor};

pub fn fast_walsh_hadamard<B: Backend>(x: Tensor<B, 2>) -> Tensor<B, 2> {
    let [n, d] = x.dims();
    let p = d.next_power_of_two();
    let dev = x.device();
    let (x_pad, was_padded) = if p != d {
        (
            Tensor::cat(vec![x, Tensor::zeros([n, p - d], &dev)], 1),
            true,
        )
    } else {
        (x, false)
    };
    let mut h = 1usize;
    let mut out = x_pad;
    while h < p {
        let step = 2 * h;
        let r = out.reshape([n, p / step, 2, h]);
        let l = r
            .clone()
            .slice([0..n, 0..(p / step), 0..1, 0..h])
            .squeeze_dim::<3>(2);
        let rt = r
            .slice([0..n, 0..(p / step), 1..2, 0..h])
            .squeeze_dim::<3>(2);
        out = Tensor::cat(
            vec![
                (l.clone() + rt.clone()).unsqueeze_dim::<4>(2),
                (l - rt).unsqueeze_dim::<4>(2),
            ],
            2,
        )
        .reshape([n, p]);
        h = step;
    }
    let out = out.div_scalar((p as f32).sqrt());
    if was_padded {
        out.slice([0..n, 0..d])
    } else {
        out
    }
}

pub fn weight_quant_ternary<B: Backend>(w: Tensor<B, 2>) -> Tensor<B, 2> {
    let scale = w.clone().abs().mean().unsqueeze_dims(&[0, 0]);
    let mean = w.clone().mean().unsqueeze_dims(&[0, 0]);
    let u = w.clone().sub(mean).sign().mul(scale);
    let base = w.clone().detach();
    w.add(u.sub(base))
}

pub fn activation_quant_8bit<B: Backend>(x: Tensor<B, 3>) -> Tensor<B, 3> {
    let [b, t, d] = x.dims();
    let flat = x.clone().reshape([b * t, d]);
    let scale = flat.clone().abs().max_dim(1).clamp_min(1e-5);
    let y = flat
        .div(scale.clone())
        .mul_scalar(127.0)
        .round()
        .clamp(-128.0, 127.0)
        .div_scalar(127.0)
        .mul(scale)
        .reshape([b, t, d]);
    let base = x.clone().detach();
    x.add(y.sub(base))
}

pub fn bitnet_v2_quantize<B: Backend>(x: Tensor<B, 3>, bits: usize) -> Tensor<B, 3> {
    if bits >= 16 {
        return x;
    }
    let [b, t, d] = x.dims();
    let flat = x.clone().reshape([b * t, d]);
    let rotated = fast_walsh_hadamard(flat);
    let deq_rot = if bits >= 8 {
        let scale = rotated.clone().abs().max_dim(1).clamp_min(1e-12);
        let q = rotated
            .clone()
            .div(scale.clone())
            .mul_scalar(127.0)
            .round()
            .clamp(-128.0, 127.0);
        fast_walsh_hadamard(q.div(scale).mul_scalar(1.0 / 127.0))
    } else {
        let scale = rotated.clone().abs().mean_dim(1).clamp_min(1e-12);
        let q = rotated
            .clone()
            .div(scale.clone())
            .mul_scalar(7.0)
            .round()
            .clamp(-8.0, 7.0);
        fast_walsh_hadamard(q.div(scale).mul_scalar(1.0 / 7.0))
    };
    let y = deq_rot.reshape([b, t, d]);
    let base = x.clone().detach();
    x.add(y.sub(base))
}

pub fn quantize_4bit<B: Backend>(x: Tensor<B, 3>) -> Tensor<B, 3> {
    bitnet_v2_quantize(x, 4)
}
pub fn quantize_8bit<B: Backend>(x: Tensor<B, 3>) -> Tensor<B, 3> {
    bitnet_v2_quantize(x, 8)
}

#[cfg(test)]
mod tests {
    use super::*;
    use burn::tensor::Distribution;
    use burn_ndarray::{NdArray, NdArrayDevice};
    type B = NdArray;
    fn dev() -> NdArrayDevice {
        NdArrayDevice::default()
    }

    #[test]
    fn hadamard_roundtrip() {
        let x = Tensor::<B, 2>::ones([4, 8], &dev());
        let h2x = fast_walsh_hadamard(fast_walsh_hadamard(x));
        let v: Vec<f32> = h2x
            .into_data()
            .bytes
            .chunks_exact(4)
            .map(|b| f32::from_le_bytes(b.try_into().unwrap()))
            .collect();
        for (i, val) in v.iter().enumerate() {
            assert!((val - 1.0).abs() < 0.1, "idx {i}: {val}");
        }
    }

    #[test]
    fn hadamard_non_power_of_two() {
        let h2x =
            fast_walsh_hadamard::<B>(fast_walsh_hadamard(Tensor::<B, 2>::ones([2, 7], &dev())));
        assert_eq!(h2x.dims(), [2, 7]);
    }

    #[test]
    fn weight_ternary_values() {
        let q = weight_quant_ternary(Tensor::<B, 2>::random(
            [4, 16],
            Distribution::Normal(0.0, 0.5),
            &dev(),
        ));
        let v: Vec<f32> = q
            .into_data()
            .bytes
            .chunks_exact(4)
            .map(|b| f32::from_le_bytes(b.try_into().unwrap()))
            .collect();
        let s = v.iter().map(|x| x.abs()).sum::<f32>() / v.len() as f32;
        for val in v {
            assert!(
                (val.abs() - s).abs() < 0.01 || val.abs() < 0.01,
                "{val} not ternary"
            );
        }
    }

    #[test]
    fn activation_8bit_shape() {
        assert_eq!(
            activation_quant_8bit(Tensor::<B, 3>::random(
                [2, 8, 64],
                Distribution::Normal(0.0, 1.0),
                &dev()
            ))
            .dims(),
            [2, 8, 64]
        );
    }

    #[test]
    fn v2_4bit_finite() {
        let q = quantize_4bit(Tensor::<B, 3>::random(
            [1, 8, 32],
            Distribution::Normal(0.0, 1.0),
            &dev(),
        ));
        assert!(q
            .into_data()
            .bytes
            .chunks_exact(4)
            .map(|b| f32::from_le_bytes(b.try_into().unwrap()))
            .all(|v| v.is_finite()));
    }

    #[test]
    fn v2_8bit_finite() {
        let q = quantize_8bit(Tensor::<B, 3>::random(
            [2, 16, 128],
            Distribution::Normal(0.0, 1.0),
            &dev(),
        ));
        assert!(q
            .into_data()
            .bytes
            .chunks_exact(4)
            .map(|b| f32::from_le_bytes(b.try_into().unwrap()))
            .all(|v| v.is_finite()));
    }

    #[test]
    fn v2_pass_through() {
        let x = Tensor::<B, 3>::random([1, 4, 16], Distribution::Normal(0.0, 1.0), &dev());
        let q = bitnet_v2_quantize(x.clone(), 16);
        let d: Vec<f32> = (x - q)
            .into_data()
            .bytes
            .chunks_exact(4)
            .map(|b| f32::from_le_bytes(b.try_into().unwrap()))
            .collect();
        assert!(d.iter().all(|&v| v.abs() < 1e-5));
    }
}