use candle_core::Tensor;
use crate::{
PlaidError,
Result,
device::default_device,
distance::squared_l2,
kmeans::{assign_as_tensor, nearest_centroid},
};
const ENCODE_CHUNK_BYTES: usize = 128 * 1024 * 1024;
fn encode_chunk_rows(dim: usize, _packed_bytes: usize) -> usize {
let bytes_per_row = dim * std::mem::size_of::<f32>();
(ENCODE_CHUNK_BYTES / bytes_per_row).max(1)
}
#[derive(Debug, Clone)]
pub struct ResidualCodec {
pub nbits: u32,
pub dim: usize,
pub centroids: Vec<f32>,
pub bucket_cutoffs: Vec<f32>,
pub bucket_weights: Vec<f32>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct EncodedVector {
pub centroid_id: u32,
pub codes: Vec<u8>,
}
pub fn packed_bytes_per_vector(dim: usize, nbits: u32) -> usize {
assert_supported_nbits(nbits);
(dim * nbits as usize).div_ceil(8)
}
fn assert_supported_nbits(nbits: u32) {
assert!(
matches!(nbits, 1 | 2 | 4 | 8),
"packed codec: nbits must be 1, 2, 4, or 8 (got {nbits})",
);
}
fn pack_codes(unpacked: &[u8], nbits: u32) -> Vec<u8> {
assert_supported_nbits(nbits);
if nbits == 8 {
return unpacked.to_vec();
}
let codes_per_byte = 8 / nbits as usize;
let mask: u8 = ((1u16 << nbits) - 1) as u8;
let n_bytes = unpacked.len().div_ceil(codes_per_byte);
let mut packed = vec![0u8; n_bytes];
for (i, &code) in unpacked.iter().enumerate() {
let byte_idx = i / codes_per_byte;
let bit_off = (i % codes_per_byte) * nbits as usize;
packed[byte_idx] |= (code & mask) << bit_off;
}
packed
}
pub fn read_code(packed: &[u8], i: usize, nbits: u32) -> u8 {
assert_supported_nbits(nbits);
if nbits == 8 {
return packed[i];
}
let codes_per_byte = 8 / nbits as usize;
let mask: u8 = ((1u16 << nbits) - 1) as u8;
let byte_idx = i / codes_per_byte;
let bit_off = (i % codes_per_byte) * nbits as usize;
(packed[byte_idx] >> bit_off) & mask
}
pub struct DecodeTable {
weights: Vec<f32>,
codes_per_byte: usize,
nbits: u32,
}
impl DecodeTable {
pub fn new(codec: &ResidualCodec) -> Self {
assert_supported_nbits(codec.nbits);
let codes_per_byte = 8 / codec.nbits as usize;
let entries = 256;
let mut weights = vec![0.0f32; entries * codes_per_byte];
let mask: u8 = ((1u16 << codec.nbits) - 1) as u8;
for b in 0u16..256 {
let byte = b as u8;
for k in 0..codes_per_byte {
let code = (byte >> (k * codec.nbits as usize)) & mask;
weights[b as usize * codes_per_byte + k] =
codec.bucket_weights[code as usize];
}
}
Self {
weights,
codes_per_byte,
nbits: codec.nbits,
}
}
pub fn weights_for(&self, byte: u8) -> &[f32] {
let start = byte as usize * self.codes_per_byte;
&self.weights[start..start + self.codes_per_byte]
}
pub fn weights_flat(&self) -> &[f32] {
&self.weights
}
pub fn codes_per_byte(&self) -> usize {
self.codes_per_byte
}
pub fn nbits(&self) -> u32 {
self.nbits
}
}
impl ResidualCodec {
pub fn num_buckets(&self) -> usize {
1usize << self.nbits
}
pub fn num_centroids(&self) -> usize {
self.centroids.len() / self.dim
}
pub fn packed_bytes(&self) -> usize {
packed_bytes_per_vector(self.dim, self.nbits)
}
pub fn validate(&self) -> Result<()> {
if self.dim == 0 {
return Err(PlaidError::InvalidCodec(
"codec: dim must be positive".into(),
));
}
if !matches!(self.nbits, 1 | 2 | 4 | 8) {
return Err(PlaidError::InvalidCodec(format!(
"codec: nbits must be 1, 2, 4, or 8, got {}",
self.nbits
)));
}
if !self.centroids.len().is_multiple_of(self.dim)
|| self.centroids.is_empty()
{
return Err(PlaidError::InvalidCodec(format!(
"codec: centroids length {} is not a positive multiple of dim {}",
self.centroids.len(),
self.dim,
)));
}
let expected_buckets = self.num_buckets();
if self.bucket_weights.len() != expected_buckets {
return Err(PlaidError::InvalidCodec(format!(
"codec: expected {} bucket_weights, got {}",
expected_buckets,
self.bucket_weights.len(),
)));
}
if self.bucket_cutoffs.len() != expected_buckets - 1 {
return Err(PlaidError::InvalidCodec(format!(
"codec: expected {} bucket_cutoffs, got {}",
expected_buckets - 1,
self.bucket_cutoffs.len(),
)));
}
for pair in self.bucket_cutoffs.windows(2) {
if pair[0] > pair[1] || pair[0].is_nan() || pair[1].is_nan() {
return Err(PlaidError::InvalidCodec(
"codec: bucket_cutoffs must be non-decreasing and finite"
.into(),
));
}
}
Ok(())
}
pub fn encode_vector(&self, vector: &[f32]) -> Result<EncodedVector> {
self.validate()?;
assert_eq!(
vector.len(),
self.dim,
"encode_vector: expected {} dims, got {}",
self.dim,
vector.len(),
);
let centroid_id = nearest_centroid(vector, &self.centroids, self.dim);
let centroid_slice = &self.centroids
[centroid_id * self.dim..(centroid_id + 1) * self.dim];
let unpacked: Vec<u8> = vector
.iter()
.zip(centroid_slice.iter())
.map(|(v, c)| bucket_for_value(*v - *c, &self.bucket_cutoffs))
.collect();
let codes = pack_codes(&unpacked, self.nbits);
Ok(EncodedVector {
centroid_id: centroid_id as u32,
codes,
})
}
pub fn batch_encode_tokens(
&self,
tokens: &[f32],
) -> Result<(Vec<u32>, Vec<u8>)> {
let chunk_rows = encode_chunk_rows(self.dim, self.packed_bytes());
self.batch_encode_tokens_with_chunk_rows(tokens, chunk_rows)
}
pub fn batch_encode_tokens_with_chunk_rows(
&self,
tokens: &[f32],
chunk_rows: usize,
) -> Result<(Vec<u32>, Vec<u8>)> {
assert!(
chunk_rows > 0,
"batch_encode_tokens_with_chunk_rows: chunk_rows must be positive"
);
self.validate()?;
assert!(
tokens.len().is_multiple_of(self.dim),
"batch_encode_tokens_with_chunk_rows: tokens length {} is not a multiple of dim {}",
tokens.len(),
self.dim,
);
let n = tokens.len() / self.dim;
if n == 0 {
return Ok((Vec::new(), Vec::new()));
}
let device = default_device();
let k = self.num_centroids();
let packed_per_token = self.packed_bytes();
let codes_per_byte = 8 / self.nbits as usize;
let centroids_dev =
Tensor::from_slice(&self.centroids, (k, self.dim), device)?;
let cutoffs_dev = Tensor::from_slice(
&self.bucket_cutoffs,
(self.bucket_cutoffs.len(),),
device,
)?;
let shift_weights: Vec<u32> = (0..codes_per_byte)
.map(|slot| 1u32 << (slot as u32 * self.nbits))
.collect();
let shifts_dev =
Tensor::from_slice(&shift_weights, (1, 1, codes_per_byte), device)?;
let mut centroid_ids: Vec<u32> = Vec::with_capacity(n);
let mut packed_codes: Vec<u8> =
Vec::with_capacity(n * packed_per_token);
let mut start = 0usize;
while start < n {
let len = chunk_rows.min(n - start);
let slice = &tokens[start * self.dim..(start + len) * self.dim];
let tile_tensor =
Tensor::from_slice(slice, (len, self.dim), device)?;
let (tile_cids, tile_codes) = self
.encode_chunk_on_tensor_with_state(
&tile_tensor,
len,
¢roids_dev,
&cutoffs_dev,
&shifts_dev,
codes_per_byte,
packed_per_token,
)?;
centroid_ids.extend(tile_cids);
packed_codes.extend(tile_codes);
start += len;
}
Ok((centroid_ids, packed_codes))
}
#[allow(clippy::too_many_arguments)]
fn encode_chunk_on_tensor_with_state(
&self,
tile: &Tensor,
len: usize,
centroids_dev: &Tensor,
cutoffs_dev: &Tensor,
shifts_dev: &Tensor,
codes_per_byte: usize,
packed_per_token: usize,
) -> Result<(Vec<u32>, Vec<u8>)> {
let device = tile.device();
let assign_chunk = assign_as_tensor(tile, centroids_dev)?;
let retrieved = centroids_dev.index_select(&assign_chunk, 0)?;
let residuals = tile.sub(&retrieved)?;
let mut buckets =
Tensor::zeros((len, self.dim), candle_core::DType::U32, device)?;
for i in 0..self.bucket_cutoffs.len() {
let cutoff = cutoffs_dev.narrow(0, i, 1)?;
let hit = residuals
.broadcast_ge(&cutoff)?
.to_dtype(candle_core::DType::U32)?;
buckets = buckets.add(&hit)?;
}
let padded_dim = packed_per_token * codes_per_byte;
let buckets_padded = if padded_dim == self.dim {
buckets
} else {
let pad_len = padded_dim - self.dim;
let pad =
Tensor::zeros((len, pad_len), candle_core::DType::U32, device)?;
Tensor::cat(&[&buckets, &pad], 1)?
};
let packed_u32 = buckets_padded
.reshape((len, packed_per_token, codes_per_byte))?
.broadcast_mul(shifts_dev)?
.sum(2)?;
let packed_u8 = packed_u32.to_dtype(candle_core::DType::U8)?;
Ok((
assign_chunk.to_vec1::<u32>()?,
packed_u8.flatten_all()?.to_vec1::<u8>()?,
))
}
pub fn batch_encode_tokens_on_tensor(
&self,
tokens_tensor: &Tensor,
tokens: &[f32],
) -> Result<(Vec<u32>, Vec<u8>)> {
self.validate()?;
assert!(
tokens.len().is_multiple_of(self.dim),
"batch_encode_tokens_on_tensor: tokens length {} is not a multiple of dim {}",
tokens.len(),
self.dim,
);
let n = tokens.len() / self.dim;
if n == 0 {
return Ok((Vec::new(), Vec::new()));
}
let device = tokens_tensor.device();
let k = self.num_centroids();
let packed_per_token = self.packed_bytes();
let codes_per_byte = 8 / self.nbits as usize;
let centroids_dev =
Tensor::from_slice(&self.centroids, (k, self.dim), device)?;
let cutoffs_dev = Tensor::from_slice(
&self.bucket_cutoffs,
(self.bucket_cutoffs.len(),),
device,
)?;
let shift_weights: Vec<u32> = (0..codes_per_byte)
.map(|slot| 1u32 << (slot as u32 * self.nbits))
.collect();
let shifts_dev =
Tensor::from_slice(&shift_weights, (1, 1, codes_per_byte), device)?;
let mut centroid_ids: Vec<u32> = Vec::with_capacity(n);
let mut packed_codes: Vec<u8> =
Vec::with_capacity(n * packed_per_token);
let chunk_rows =
encode_chunk_rows(self.dim, packed_per_token).min(n).max(1);
let mut start = 0usize;
while start < n {
let len = chunk_rows.min(n - start);
let tile = tokens_tensor.narrow(0, start, len)?;
let (tile_cids, tile_codes) = self
.encode_chunk_on_tensor_with_state(
&tile,
len,
¢roids_dev,
&cutoffs_dev,
&shifts_dev,
codes_per_byte,
packed_per_token,
)?;
centroid_ids.extend(tile_cids);
packed_codes.extend(tile_codes);
start += len;
}
Ok((centroid_ids, packed_codes))
}
pub fn decode_vector(&self, encoded: &EncodedVector) -> Result<Vec<f32>> {
let table = DecodeTable::new(self);
self.decode_vector_with_table(encoded, &table)
}
pub fn decode_vector_with_table(
&self,
encoded: &EncodedVector,
table: &DecodeTable,
) -> Result<Vec<f32>> {
self.validate()?;
assert_eq!(
table.nbits, self.nbits,
"decode_vector_with_table: table nbits {} != codec nbits {}",
table.nbits, self.nbits,
);
let expected_bytes = self.packed_bytes();
assert_eq!(
encoded.codes.len(),
expected_bytes,
"decode_vector_with_table: expected {expected_bytes} packed bytes, got {}",
encoded.codes.len(),
);
let centroid_id = encoded.centroid_id as usize;
assert!(
centroid_id < self.num_centroids(),
"decode_vector_with_table: centroid_id {} out of range 0..{}",
centroid_id,
self.num_centroids(),
);
let centroid_slice = &self.centroids
[centroid_id * self.dim..(centroid_id + 1) * self.dim];
let codes_per_byte = table.codes_per_byte;
let mut out = Vec::with_capacity(self.dim);
for (byte_idx, &byte) in encoded.codes.iter().enumerate() {
let weights = table.weights_for(byte);
let base_dim = byte_idx * codes_per_byte;
for (k, &w) in weights.iter().enumerate() {
let dim_idx = base_dim + k;
if dim_idx >= self.dim {
break;
}
out.push(centroid_slice[dim_idx] + w);
}
}
Ok(out)
}
pub fn reconstruction_error(&self, vector: &[f32]) -> Result<f32> {
let encoded = self.encode_vector(vector)?;
let decoded = self.decode_vector(&encoded)?;
Ok(squared_l2(vector, &decoded))
}
}
pub fn train_quantizer(
mut residuals: Vec<f32>,
nbits: u32,
) -> (Vec<f32>, Vec<f32>) {
assert!(!residuals.is_empty(), "train_quantizer: empty sample");
assert!(
nbits > 0 && nbits <= 8,
"train_quantizer: nbits must be in 1..=8, got {nbits}"
);
assert!(
residuals.iter().all(|v| !v.is_nan()),
"train_quantizer: residual sample contains NaN"
);
let num_buckets = 1usize << nbits;
let n = residuals.len();
residuals.sort_unstable_by(|a, b| a.total_cmp(b));
let bucket_bounds = |i: usize| -> (usize, usize) {
let start = i * n / num_buckets;
let end = if i + 1 == num_buckets {
n
} else {
(i + 1) * n / num_buckets
};
(start, end)
};
let cutoffs: Vec<f32> = (1..num_buckets)
.map(|i| residuals[i * n / num_buckets])
.collect();
let weights: Vec<f32> = (0..num_buckets)
.map(|i| {
let (start, end) = bucket_bounds(i);
if start == end {
let idx = start.min(n - 1);
residuals[idx]
} else {
let slice = &residuals[start..end];
slice.iter().sum::<f32>() / slice.len() as f32
}
})
.collect();
(cutoffs, weights)
}
fn bucket_for_value(value: f32, cutoffs: &[f32]) -> u8 {
let mut idx = 0u8;
for cutoff in cutoffs {
if value >= *cutoff {
idx += 1;
} else {
break;
}
}
idx
}
#[cfg(test)]
mod tests {
use super::*;
fn two_bit_1d_codec_with_centroids(centroids: Vec<f32>) -> ResidualCodec {
ResidualCodec {
nbits: 2,
dim: 1,
centroids,
bucket_cutoffs: vec![-0.5, 0.0, 0.5],
bucket_weights: vec![-0.75, -0.25, 0.25, 0.75],
}
}
#[test]
fn decode_with_lookup_table_matches_scalar_decode() {
for &nbits in &[1u32, 2, 4, 8] {
let num_buckets = 1usize << nbits;
let codec = ResidualCodec {
nbits,
dim: 16,
centroids: (0..16).map(|i| i as f32 * 0.1).collect(),
bucket_cutoffs: (1..num_buckets)
.map(|i| (i as f32 / num_buckets as f32) - 0.5)
.collect(),
bucket_weights: (0..num_buckets)
.map(|i| (i as f32 + 0.5) / num_buckets as f32 - 0.5)
.collect(),
};
let input: Vec<f32> =
(0..16).map(|i| i as f32 * 0.05 - 0.3).collect();
let encoded = codec.encode_vector(&input).unwrap();
let scalar = codec.decode_vector(&encoded).unwrap();
let table = DecodeTable::new(&codec);
let via_table =
codec.decode_vector_with_table(&encoded, &table).unwrap();
assert_eq!(scalar, via_table, "mismatch at nbits={nbits}");
}
}
#[test]
fn pack_then_read_code_recovers_every_input() {
for &nbits in &[1u32, 2, 4, 8] {
let num_buckets = 1usize << nbits;
let unpacked: Vec<u8> = (0..32u8)
.map(|i| (i as usize % num_buckets) as u8)
.collect();
let packed = pack_codes(&unpacked, nbits);
for (i, &expected) in unpacked.iter().enumerate() {
let got = read_code(&packed, i, nbits);
assert_eq!(
got, expected,
"nbits={nbits} position {i}: got {got}, expected {expected}",
);
}
assert_eq!(
packed.len(),
packed_bytes_per_vector(unpacked.len(), nbits),
);
}
}
#[test]
fn encode_vector_produces_packed_codes_at_two_bits() {
let codec = ResidualCodec {
nbits: 2,
dim: 8,
centroids: vec![0.0; 8],
bucket_cutoffs: vec![-0.5, 0.0, 0.5],
bucket_weights: vec![-0.75, -0.25, 0.25, 0.75],
};
let encoded = codec.encode_vector(&[0.1f32; 8]).unwrap();
assert_eq!(encoded.codes.len(), 2);
}
#[test]
fn encode_vector_produces_packed_codes_at_four_bits() {
let codec = ResidualCodec {
nbits: 4,
dim: 8,
centroids: vec![0.0; 8],
bucket_cutoffs: (0..15).map(|i| i as f32 / 15.0 - 0.5).collect(),
bucket_weights: (0..16).map(|i| i as f32 / 16.0 - 0.5).collect(),
};
let encoded = codec.encode_vector(&[0.1f32; 8]).unwrap();
assert_eq!(encoded.codes.len(), 4);
}
#[test]
fn encode_decode_roundtrip_at_every_supported_nbits() {
for nbits in [1u32, 2, 4, 8] {
let num_buckets = 1usize << nbits;
let bucket_cutoffs: Vec<f32> = (1..num_buckets)
.map(|i| (i as f32 / num_buckets as f32) - 0.5)
.collect();
let bucket_weights: Vec<f32> = (0..num_buckets)
.map(|i| (i as f32 + 0.5) / num_buckets as f32 - 0.5)
.collect();
let codec = ResidualCodec {
nbits,
dim: 8,
centroids: vec![0.0; 8],
bucket_cutoffs,
bucket_weights,
};
let input = [-0.4f32, -0.1, 0.0, 0.25, 0.49, -0.25, 0.1, 0.3];
let encoded = codec.encode_vector(&input).unwrap();
let decoded = codec.decode_vector(&encoded).unwrap();
let max_err = input
.iter()
.zip(decoded.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
let tolerance = 1.0 / num_buckets as f32;
assert!(
max_err <= tolerance,
"nbits={nbits}: max_err={max_err}, tolerance={tolerance}",
);
}
}
#[test]
fn bucket_for_value_places_below_first_cutoff_in_bucket_zero() {
let cutoffs = [-0.5, 0.0, 0.5];
assert_eq!(bucket_for_value(-1.0, &cutoffs), 0);
}
#[test]
fn bucket_for_value_places_at_or_above_last_cutoff_in_top_bucket() {
let cutoffs = [-0.5, 0.0, 0.5];
assert_eq!(bucket_for_value(0.5, &cutoffs), 3);
assert_eq!(bucket_for_value(9.9, &cutoffs), 3);
}
#[test]
fn bucket_for_value_picks_intermediate_buckets() {
let cutoffs = [-0.5, 0.0, 0.5];
assert_eq!(bucket_for_value(-0.25, &cutoffs), 1);
assert_eq!(bucket_for_value(0.25, &cutoffs), 2);
}
#[test]
fn num_buckets_is_two_to_the_nbits() {
let codec = two_bit_1d_codec_with_centroids(vec![0.0]);
assert_eq!(codec.num_buckets(), 4);
let mut four_bit = codec.clone();
four_bit.nbits = 4;
four_bit.bucket_cutoffs = (0..15).map(|i| i as f32 / 15.0).collect();
four_bit.bucket_weights = (0..16).map(|i| i as f32).collect();
assert_eq!(four_bit.num_buckets(), 16);
}
#[test]
fn encode_picks_nearest_centroid() {
let codec = two_bit_1d_codec_with_centroids(vec![0.0, 10.0]);
let encoded = codec.encode_vector(&[9.0]).unwrap();
assert_eq!(encoded.centroid_id, 1);
}
#[test]
fn decode_inverts_a_known_encoding() {
let codec = two_bit_1d_codec_with_centroids(vec![0.0]);
let encoded = codec.encode_vector(&[-0.3]).unwrap();
assert_eq!(encoded.codes, vec![1]);
let decoded = codec.decode_vector(&encoded).unwrap();
assert_eq!(decoded, vec![-0.25]);
}
#[test]
fn encode_then_decode_stays_inside_bucket_half_width() {
let codec = two_bit_1d_codec_with_centroids(vec![0.0]);
for &value in &[-0.4f32, -0.1, 0.0, 0.2, 0.4] {
let encoded = codec.encode_vector(&[value]).unwrap();
let decoded = codec.decode_vector(&encoded).unwrap();
assert!(
(decoded[0] - value).abs() <= 0.25,
"value {value} -> decoded {d}",
d = decoded[0],
);
}
}
#[test]
fn reconstruction_error_is_zero_when_residual_exactly_matches_weight() {
let codec = two_bit_1d_codec_with_centroids(vec![0.0]);
assert_eq!(codec.reconstruction_error(&[-0.25]).unwrap(), 0.0);
}
#[test]
fn validate_rejects_wrong_number_of_cutoffs() {
let mut codec = two_bit_1d_codec_with_centroids(vec![0.0]);
codec.bucket_cutoffs.push(1.0); assert!(codec.validate().is_err());
}
#[test]
fn validate_rejects_non_monotonic_cutoffs() {
let mut codec = two_bit_1d_codec_with_centroids(vec![0.0]);
codec.bucket_cutoffs = vec![0.5, 0.0, 0.5];
assert!(codec.validate().is_err());
}
#[test]
#[should_panic(expected = "packed bytes")]
fn decode_panics_on_wrong_packed_code_length() {
let codec = two_bit_1d_codec_with_centroids(vec![0.0]);
let bad = EncodedVector {
centroid_id: 0,
codes: vec![0, 0],
};
let _ = codec.decode_vector(&bad).unwrap();
}
#[test]
fn train_quantizer_produces_right_number_of_cutoffs_and_weights() {
let residuals: Vec<f32> =
(0..1000).map(|i| i as f32 / 1000.0).collect();
let (cutoffs, weights) = train_quantizer(residuals, 2);
assert_eq!(cutoffs.len(), 3);
assert_eq!(weights.len(), 4);
}
#[test]
fn train_quantizer_cutoffs_are_monotonic() {
let residuals: Vec<f32> =
(0..2048).map(|i| (i as f32 / 2048.0) - 0.5).collect();
let (cutoffs, _) = train_quantizer(residuals, 4);
for pair in cutoffs.windows(2) {
assert!(
pair[0] <= pair[1],
"cutoffs must be non-decreasing: {pair:?}"
);
}
}
#[test]
fn train_quantizer_on_uniform_data_gives_quartile_cutoffs() {
let residuals: Vec<f32> = (0..1000).map(|i| i as f32).collect();
let (cutoffs, _) = train_quantizer(residuals, 2);
assert!((cutoffs[0] - 250.0).abs() < 1.0);
assert!((cutoffs[1] - 500.0).abs() < 1.0);
assert!((cutoffs[2] - 750.0).abs() < 1.0);
}
#[test]
fn train_quantizer_weights_bracket_cutoffs() {
let residuals: Vec<f32> = (0..1024).map(|i| i as f32).collect();
let (cutoffs, weights) = train_quantizer(residuals, 2);
assert!(weights[0] < cutoffs[0]);
assert!(weights[3] > cutoffs[2]);
assert!(cutoffs[0] <= weights[1] && weights[1] < cutoffs[1]);
assert!(cutoffs[1] <= weights[2] && weights[2] < cutoffs[2]);
}
#[test]
#[should_panic(expected = "empty sample")]
fn train_quantizer_panics_on_empty_sample() {
let _ = train_quantizer(Vec::new(), 2);
}
#[test]
#[should_panic(expected = "NaN")]
fn train_quantizer_panics_on_nan() {
let _ = train_quantizer(vec![0.1, f32::NAN, 0.3], 2);
}
#[test]
fn trained_codec_round_trips_within_reasonable_error() {
let training: Vec<f32> =
(0..2048).map(|i| (i as f32 / 2048.0) - 0.5).collect();
let (cutoffs, weights) = train_quantizer(training, 4);
let codec = ResidualCodec {
nbits: 4,
dim: 1,
centroids: vec![0.0],
bucket_cutoffs: cutoffs,
bucket_weights: weights,
};
codec.validate().unwrap();
let mut max_err: f32 = 0.0;
for v in &[-0.4f32, -0.1, 0.0, 0.25, 0.49] {
let err = codec.reconstruction_error(&[*v]).unwrap().sqrt();
max_err = max_err.max(err);
}
assert!(
max_err < 0.05,
"max reconstruction error {max_err} above tolerance"
);
}
#[test]
fn encode_and_decode_roundtrip_multi_dim_stays_close() {
let codec = ResidualCodec {
nbits: 2,
dim: 2,
centroids: vec![1.0, 1.0],
bucket_cutoffs: vec![-0.5, 0.0, 0.5],
bucket_weights: vec![-0.75, -0.25, 0.25, 0.75],
};
let input = [1.1f32, 0.7];
let encoded = codec.encode_vector(&input).unwrap();
let decoded = codec.decode_vector(&encoded).unwrap();
for (d, i) in decoded.iter().zip(input.iter()) {
assert!((d - i).abs() <= 0.25);
}
}
}