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));
}
}