use std::{
collections::{HashMap, HashSet},
path::Path,
};
use distances::{number::UInt, Number};
use crate::{Cluster, Dataset, Instance};
use super::{DecoderFn, EncoderFn, SquishyBall};
#[derive(Debug)]
pub struct CodecData<I: Instance, U: UInt, M: Instance> {
root: SquishyBall<U>,
centers: HashMap<usize, I>,
encoder: EncoderFn<I>,
leaf_data: LeafData<I>,
metric: fn(&I, &I) -> U,
is_expensive: bool,
metadata: Vec<M>,
permuted_indices: Vec<usize>,
}
impl<I: Instance, U: UInt, M: Instance> CodecData<I, U, M> {
pub fn new<D: Dataset<I, U>>(
mut root: SquishyBall<U>,
data: &D,
encoder: EncoderFn<I>,
decoder: DecoderFn<I>,
metadata: Vec<M>,
) -> Result<Self, String> {
let permuted_indices = data
.permuted_indices()
.map_or_else(|| (0..data.cardinality()).collect(), <[usize]>::to_vec);
root.trim();
let subtree = root.compressible_subtree();
let centers = subtree.iter().map(|c| c.arg_center()).collect::<HashSet<_>>();
let centers = centers
.into_iter()
.map(|i| (i, data[i].clone()))
.collect::<HashMap<_, _>>();
let mut bytes = Vec::new();
for leaf in root.compressible_leaves_mut().into_iter().filter(|c| c.squish()) {
leaf.set_codec_offset(bytes.len());
let center = &data[leaf.arg_center()];
let encodings = leaf
.indices()
.map(|i| encoder(center, &data[i]))
.collect::<Result<Vec<_>, _>>()?;
bytes.extend_from_slice(&leaf.cardinality().to_le_bytes());
for encoding in encodings {
let len = encoding.len();
bytes.extend_from_slice(&len.to_le_bytes());
bytes.extend_from_slice(&encoding);
}
}
let bytes = bytes.into_boxed_slice();
let leaf_data = LeafData { bytes, decoder };
Ok(Self {
root,
centers,
encoder,
leaf_data,
metric: data.metric(),
is_expensive: data.is_metric_expensive(),
metadata,
permuted_indices,
})
}
pub fn load_leaf_data(&self, leaf: &SquishyBall<U>) -> Result<Vec<I>, String> {
let offset = leaf.codec_offset().ok_or("Leaf has no codec offset")?;
let center = &self.centers[&leaf.arg_center()];
self.leaf_data.load_leaf(center, offset)
}
pub const fn root(&self) -> &SquishyBall<U> {
&self.root
}
pub const fn centers(&self) -> &HashMap<usize, I> {
&self.centers
}
pub fn metadata(&self) -> &[M] {
&self.metadata
}
pub fn permuted_indices(&self) -> &[usize] {
&self.permuted_indices
}
pub fn metric(&self) -> fn(&I, &I) -> U {
self.metric
}
pub const fn is_expensive(&self) -> bool {
self.is_expensive
}
pub fn save(&self, path: &Path) -> Result<(), String> {
if let Some(parent) = path.parent() {
if !parent.exists() {
return Err(format!("Parent directory does not exist: {parent:?}"));
}
} else {
return Err("Path has no parent directory".to_string());
}
if path.exists() {
std::fs::remove_dir_all(path).map_err(|e| e.to_string())?;
}
std::fs::create_dir(path).map_err(|e| e.to_string())?;
let root_path = path.join("root.bin");
self.root.save(&root_path)?;
let centers_path = path.join("centers.bin");
let centers = encode_centers(&self.root, &self.centers, self.encoder)?;
std::fs::write(centers_path, ¢ers).map_err(|e| e.to_string())?;
let leaf_data_path = path.join("leaf_data.bin");
std::fs::write(leaf_data_path, &self.leaf_data.bytes).map_err(|e| e.to_string())?;
let metadata_path = path.join("metadata.bin");
let metadata = self.metadata.iter().map(M::to_bytes).collect::<Vec<_>>();
let metadata = bincode::serialize(&metadata).map_err(|e| e.to_string())?;
std::fs::write(metadata_path, metadata).map_err(|e| e.to_string())?;
let permuted_indices_path = path.join("permuted_indices.bin");
let permuted_indices = bincode::serialize(&self.permuted_indices).map_err(|e| e.to_string())?;
std::fs::write(permuted_indices_path, permuted_indices).map_err(|e| e.to_string())?;
Ok(())
}
pub fn load(
path: &Path,
metric: fn(&I, &I) -> U,
is_expensive: bool,
encoder: EncoderFn<I>,
decoder: DecoderFn<I>,
) -> Result<Self, String> {
if !path.exists() {
return Err(format!("Directory does not exist: {path:?}"));
}
if !path.is_dir() {
return Err(format!("Path is not a directory: {path:?}"));
}
let root_path = path.join("root.bin");
let centers_path = path.join("centers.bin");
let leaf_data_path = path.join("leaf_data.bin");
let metadata_path = path.join("metadata.bin");
let permuted_indices_path = path.join("permuted_indices.bin");
for file in [
&root_path,
¢ers_path,
&leaf_data_path,
&metadata_path,
&permuted_indices_path,
] {
if !file.exists() {
return Err(format!("File does not exist: {file:?}"));
}
}
let root = SquishyBall::load(&root_path)?;
let centers = std::fs::read(¢ers_path).map_err(|e| e.to_string())?;
let centers = decode_centers(&root, ¢ers, decoder)?;
let leaf_data = std::fs::read(&leaf_data_path).map_err(|e| e.to_string())?;
let leaf_data = LeafData {
bytes: leaf_data.into_boxed_slice(),
decoder,
};
let metadata = std::fs::read(&metadata_path).map_err(|e| e.to_string())?;
let metadata: Vec<Vec<u8>> = bincode::deserialize(&metadata).map_err(|e| e.to_string())?;
let metadata = metadata
.into_iter()
.map(|m| M::from_bytes(&m))
.collect::<Result<Vec<_>, _>>()?;
let permuted_indices = std::fs::read(&permuted_indices_path).map_err(|e| e.to_string())?;
let permuted_indices = bincode::deserialize(&permuted_indices).map_err(|e| e.to_string())?;
Ok(Self {
root,
centers,
encoder,
leaf_data,
metric,
is_expensive,
metadata,
permuted_indices,
})
}
}
fn encode_centers<I: Instance, U: UInt>(
root: &SquishyBall<U>,
centers: &HashMap<usize, I>,
encoder: EncoderFn<I>,
) -> Result<Box<[u8]>, String> {
let mut bytes = Vec::new();
let root_center = centers[&root.arg_center()].to_bytes();
bytes.extend_from_slice(&root_center.len().to_le_bytes());
bytes.extend_from_slice(&root_center);
for (reference, target) in index_pairs(root) {
bytes.extend_from_slice(&reference.to_le_bytes());
bytes.extend_from_slice(&target.to_le_bytes());
let encoding = encoder(¢ers[&reference], ¢ers[&target])?;
bytes.extend_from_slice(&encoding.len().to_le_bytes());
bytes.extend_from_slice(&encoding);
}
Ok(bytes.into_boxed_slice())
}
fn index_pairs<U: UInt>(c: &SquishyBall<U>) -> Vec<(usize, usize)> {
let mut pairs = Vec::new();
if !c.squish() {
if let Some([left, right]) = c.children() {
pairs.push((c.arg_center(), left.arg_center()));
pairs.push((c.arg_center(), right.arg_center()));
pairs.append(&mut index_pairs(left));
pairs.append(&mut index_pairs(right));
}
}
pairs
}
fn decode_centers<I: Instance, U: UInt>(
root: &SquishyBall<U>,
bytes: &[u8],
decoder: DecoderFn<I>,
) -> Result<HashMap<usize, I>, String> {
let mut centers = HashMap::new();
let mut offset = 0;
let len = <usize as Number>::from_le_bytes(&bytes[offset..(offset + usize::num_bytes())]);
offset += usize::num_bytes();
let root_center = I::from_bytes(&bytes[offset..(offset + len)])?;
centers.insert(root.arg_center(), root_center);
offset += len;
while offset < bytes.len() {
let reference = <usize as Number>::from_le_bytes(&bytes[offset..(offset + usize::num_bytes())]);
offset += usize::num_bytes();
let target = <usize as Number>::from_le_bytes(&bytes[offset..(offset + usize::num_bytes())]);
offset += usize::num_bytes();
let len = <usize as Number>::from_le_bytes(&bytes[offset..(offset + usize::num_bytes())]);
offset += usize::num_bytes();
let encoding = decoder(¢ers[&reference], &bytes[offset..(offset + len)])?;
centers.insert(target, encoding);
offset += len;
}
Ok(centers)
}
#[derive(Debug)]
struct LeafData<I: Instance> {
pub bytes: Box<[u8]>,
pub decoder: DecoderFn<I>,
}
impl<I: Instance> LeafData<I> {
fn load_leaf(&self, center: &I, offset: usize) -> Result<Vec<I>, String> {
let cardinality = {
let bytes = &self.bytes[offset..(offset + usize::num_bytes())];
<usize as Number>::from_le_bytes(bytes)
};
let mut data = Vec::with_capacity(cardinality);
let mut offset = offset + usize::num_bytes();
for _ in 0..cardinality {
let len = {
let bytes = &self.bytes[offset..(offset + usize::num_bytes())];
<usize as Number>::from_le_bytes(bytes)
};
offset += usize::num_bytes();
let target = {
let bytes = &self.bytes[offset..(offset + len)];
offset += len;
(self.decoder)(center, bytes)?
};
data.push(target);
}
Ok(data)
}
}