998 lines
33 KiB
Rust
998 lines
33 KiB
Rust
//! Vector Quantization with K-means clustering
|
|
//!
|
|
//! This module provides a complete K-means implementation for vector quantization
|
|
//! including K-means++ initialization, multiple distance metrics, and online updates.
|
|
|
|
use crate::error::{CompressionError, QuantizationError, Result};
|
|
use rtx_tensor::{Device, Tensor};
|
|
|
|
/// Method for initializing the codebook
|
|
#[derive(Debug, Clone, PartialEq)]
|
|
pub enum CodebookInitialization {
|
|
/// Random initialization from data samples
|
|
Random,
|
|
/// K-means++ for better initial centroids
|
|
KMeansPlusPlus,
|
|
/// Initialize from provided data directly
|
|
FromData,
|
|
}
|
|
|
|
/// Distance metric for comparing vectors
|
|
#[derive(Debug, Clone, PartialEq)]
|
|
pub enum DistanceMetric {
|
|
/// Euclidean (L2) distance
|
|
Euclidean,
|
|
/// Cosine distance (1 - cosine similarity)
|
|
Cosine,
|
|
/// Manhattan (L1) distance
|
|
Manhattan,
|
|
}
|
|
|
|
/// Configuration for vector quantization
|
|
#[derive(Debug, Clone)]
|
|
pub struct VQConfig {
|
|
/// Number of codewords in the codebook
|
|
pub codebook_size: usize,
|
|
/// Dimension of each vector
|
|
pub vector_dim: usize,
|
|
/// Maximum iterations for K-means
|
|
pub max_iterations: usize,
|
|
/// Convergence tolerance
|
|
pub tolerance: f64,
|
|
/// Initialization method
|
|
pub initialization: CodebookInitialization,
|
|
/// Distance metric to use
|
|
pub distance_metric: DistanceMetric,
|
|
}
|
|
|
|
impl Default for VQConfig {
|
|
fn default() -> Self {
|
|
Self {
|
|
codebook_size: 256,
|
|
vector_dim: 64,
|
|
max_iterations: 100,
|
|
tolerance: 1e-6,
|
|
initialization: CodebookInitialization::KMeansPlusPlus,
|
|
distance_metric: DistanceMetric::Euclidean,
|
|
}
|
|
}
|
|
}
|
|
|
|
/// Statistics from K-means training
|
|
#[derive(Debug, Clone)]
|
|
pub struct TrainingStats {
|
|
/// Number of iterations run
|
|
pub iterations: usize,
|
|
/// Final inertia (sum of squared distances to centroids)
|
|
pub final_inertia: f64,
|
|
/// Initial inertia
|
|
pub initial_inertia: f64,
|
|
/// Whether convergence was reached
|
|
pub converged: bool,
|
|
/// Usage count per codeword
|
|
pub codeword_usage: Vec<usize>,
|
|
}
|
|
|
|
/// Vector quantizer using K-means clustering
|
|
pub struct VectorQuantizer {
|
|
config: VQConfig,
|
|
/// Codebook stored as flattened vector [codebook_size * vector_dim]
|
|
codebook_data: Option<Vec<f32>>,
|
|
codebook: Option<Tensor>,
|
|
device: Option<Device>,
|
|
online_mode: bool,
|
|
learning_rate: f64,
|
|
training_stats: Option<TrainingStats>,
|
|
}
|
|
|
|
impl VectorQuantizer {
|
|
/// Create a new vector quantizer with the given configuration
|
|
pub fn new(config: VQConfig) -> Self {
|
|
Self {
|
|
config,
|
|
codebook_data: None,
|
|
codebook: None,
|
|
device: None,
|
|
online_mode: false,
|
|
learning_rate: 0.01,
|
|
training_stats: None,
|
|
}
|
|
}
|
|
|
|
/// Fit the quantizer to data using K-means clustering
|
|
pub fn fit(&mut self, data: &Tensor) -> Result<()> {
|
|
self.device = Some(data.device().clone());
|
|
|
|
let shape = data.shape();
|
|
if shape.dims().len() != 2 {
|
|
return Err(CompressionError::Quantization(
|
|
QuantizationError::SerializationFailed(
|
|
"Data must be 2D: [num_samples, vector_dim]".to_string(),
|
|
),
|
|
));
|
|
}
|
|
|
|
let num_samples = shape.dims()[0];
|
|
let vector_dim = shape.dims()[1];
|
|
|
|
if vector_dim != self.config.vector_dim {
|
|
return Err(CompressionError::Quantization(
|
|
QuantizationError::SerializationFailed(format!(
|
|
"Vector dimension mismatch: expected {}, got {}",
|
|
self.config.vector_dim, vector_dim
|
|
)),
|
|
));
|
|
}
|
|
|
|
if num_samples < self.config.codebook_size {
|
|
return Err(CompressionError::Quantization(
|
|
QuantizationError::SerializationFailed(format!(
|
|
"Not enough samples ({}) for codebook size ({})",
|
|
num_samples, self.config.codebook_size
|
|
)),
|
|
));
|
|
}
|
|
|
|
// Extract data from tensor
|
|
let data_vec = self.extract_tensor_data(data)?;
|
|
|
|
// Run K-means
|
|
let (codebook_data, stats) = self.kmeans_fit(&data_vec, num_samples, vector_dim)?;
|
|
|
|
self.codebook_data = Some(codebook_data.clone());
|
|
self.training_stats = Some(stats);
|
|
|
|
// Create tensor from codebook data
|
|
self.codebook = Some(self.create_codebook_tensor(&codebook_data)?);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Encode data vectors to their nearest codeword indices
|
|
pub fn encode(&self, data: &Tensor) -> Result<Tensor> {
|
|
let codebook_data = self.codebook_data.as_ref().ok_or_else(|| {
|
|
CompressionError::Quantization(QuantizationError::SerializationFailed(
|
|
"Quantizer not trained. Call fit() first".to_string(),
|
|
))
|
|
})?;
|
|
|
|
let device = self.device.as_ref().unwrap();
|
|
let shape = data.shape();
|
|
let num_samples = shape.dims()[0];
|
|
let vector_dim = shape.dims()[1];
|
|
|
|
// Extract data from tensor
|
|
let data_vec = self.extract_tensor_data(data)?;
|
|
|
|
// Find nearest centroid for each sample
|
|
let mut codes = Vec::with_capacity(num_samples);
|
|
for i in 0..num_samples {
|
|
let offset = i * vector_dim;
|
|
let sample = &data_vec[offset..offset + vector_dim];
|
|
let (nearest_idx, _) = self.find_nearest_centroid(sample, codebook_data);
|
|
codes.push(nearest_idx as f32);
|
|
}
|
|
|
|
// Create tensor from codes
|
|
let codes_tensor = self.create_codes_tensor(&codes, device)?;
|
|
Ok(codes_tensor)
|
|
}
|
|
|
|
/// Decode codeword indices back to vectors
|
|
pub fn decode(&self, codes: &Tensor) -> Result<Tensor> {
|
|
let codebook_data = self.codebook_data.as_ref().ok_or_else(|| {
|
|
CompressionError::Quantization(QuantizationError::SerializationFailed(
|
|
"Quantizer not trained. Call fit() first".to_string(),
|
|
))
|
|
})?;
|
|
|
|
let device = self.device.as_ref().unwrap();
|
|
let num_samples = codes.shape().dims()[0];
|
|
let vector_dim = self.config.vector_dim;
|
|
|
|
// Extract codes from tensor
|
|
let codes_vec = self.extract_tensor_data(codes)?;
|
|
|
|
// Reconstruct vectors from codebook
|
|
let mut reconstructed = Vec::with_capacity(num_samples * vector_dim);
|
|
for code in &codes_vec {
|
|
let idx = (*code as usize).min(self.config.codebook_size - 1);
|
|
let offset = idx * vector_dim;
|
|
reconstructed.extend_from_slice(&codebook_data[offset..offset + vector_dim]);
|
|
}
|
|
|
|
// Create tensor from reconstructed data
|
|
let reconstructed_tensor =
|
|
self.create_reconstructed_tensor(&reconstructed, num_samples, device)?;
|
|
Ok(reconstructed_tensor)
|
|
}
|
|
|
|
/// Get the codebook tensor
|
|
pub fn codebook(&self) -> &Tensor {
|
|
self.codebook.as_ref().expect("Codebook not initialized")
|
|
}
|
|
|
|
/// Get the codebook size
|
|
pub fn codebook_size(&self) -> usize {
|
|
self.config.codebook_size
|
|
}
|
|
|
|
/// Get configuration
|
|
pub fn config(&self) -> &VQConfig {
|
|
&self.config
|
|
}
|
|
|
|
/// Get training statistics
|
|
pub fn training_stats(&self) -> Option<&TrainingStats> {
|
|
self.training_stats.as_ref()
|
|
}
|
|
|
|
/// Enable online updates for streaming data
|
|
pub fn enable_online_updates(&mut self, enable: bool, learning_rate: f64) {
|
|
self.online_mode = enable;
|
|
self.learning_rate = learning_rate;
|
|
}
|
|
|
|
/// Update codebook with new data (online learning)
|
|
pub fn update_online(&mut self, new_data: &Tensor) -> Result<()> {
|
|
if !self.online_mode {
|
|
return Err(CompressionError::Quantization(
|
|
QuantizationError::SerializationFailed("Online updates not enabled".to_string()),
|
|
));
|
|
}
|
|
|
|
if self.codebook_data.is_none() {
|
|
return Err(CompressionError::Quantization(
|
|
QuantizationError::SerializationFailed("Quantizer not trained".to_string()),
|
|
));
|
|
}
|
|
|
|
let data_vec = self.extract_tensor_data(new_data)?;
|
|
let num_samples = new_data.shape().dims()[0];
|
|
let vector_dim = self.config.vector_dim;
|
|
let learning_rate = self.learning_rate as f32;
|
|
let codebook_size = self.config.codebook_size;
|
|
|
|
// Get mutable reference to codebook data
|
|
let codebook_data = self.codebook_data.as_mut().unwrap();
|
|
|
|
// Online update: move centroids toward new data
|
|
for i in 0..num_samples {
|
|
let offset = i * vector_dim;
|
|
let sample = &data_vec[offset..offset + vector_dim];
|
|
|
|
// Find nearest centroid inline to avoid borrow issues
|
|
let mut nearest_idx = 0;
|
|
let mut min_dist = f32::INFINITY;
|
|
for c in 0..codebook_size {
|
|
let centroid_offset = c * vector_dim;
|
|
let mut dist = 0.0f32;
|
|
for j in 0..vector_dim {
|
|
let diff = sample[j] - codebook_data[centroid_offset + j];
|
|
dist += diff * diff;
|
|
}
|
|
if dist < min_dist {
|
|
min_dist = dist;
|
|
nearest_idx = c;
|
|
}
|
|
}
|
|
|
|
// Update centroid: c = c + lr * (x - c)
|
|
let centroid_offset = nearest_idx * vector_dim;
|
|
for j in 0..vector_dim {
|
|
let diff = sample[j] - codebook_data[centroid_offset + j];
|
|
codebook_data[centroid_offset + j] += learning_rate * diff;
|
|
}
|
|
}
|
|
|
|
// Update tensor codebook
|
|
let device = Device::try_default()?;
|
|
let tensor = Tensor::from_slice(codebook_data, &[codebook_size, vector_dim], &device)?;
|
|
self.codebook = Some(tensor);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Analyze codebook usage statistics
|
|
pub fn analyze_codebook_usage(&self, codes: &Tensor) -> Result<Vec<f64>> {
|
|
let codes_vec = self.extract_tensor_data(codes)?;
|
|
let total_codes = codes_vec.len();
|
|
|
|
let mut usage_counts = vec![0usize; self.config.codebook_size];
|
|
for &code in &codes_vec {
|
|
let idx = (code as usize).min(self.config.codebook_size - 1);
|
|
usage_counts[idx] += 1;
|
|
}
|
|
|
|
// Convert to proportions
|
|
let usage_stats: Vec<f64> = usage_counts
|
|
.iter()
|
|
.map(|&count| count as f64 / total_codes as f64)
|
|
.collect();
|
|
|
|
Ok(usage_stats)
|
|
}
|
|
|
|
/// Prune unused or underused codewords
|
|
pub fn prune_codebook(&self, min_usage_threshold: f64) -> Result<Self> {
|
|
let codebook_data = self.codebook_data.as_ref().ok_or_else(|| {
|
|
CompressionError::Quantization(QuantizationError::SerializationFailed(
|
|
"Quantizer not trained".to_string(),
|
|
))
|
|
})?;
|
|
|
|
// Get usage from training stats if available
|
|
let usage_stats = if let Some(stats) = &self.training_stats {
|
|
let total: usize = stats.codeword_usage.iter().sum();
|
|
stats
|
|
.codeword_usage
|
|
.iter()
|
|
.map(|&c| c as f64 / total.max(1) as f64)
|
|
.collect::<Vec<_>>()
|
|
} else {
|
|
vec![1.0 / self.config.codebook_size as f64; self.config.codebook_size]
|
|
};
|
|
|
|
// Select codewords that meet threshold
|
|
let active_indices: Vec<usize> = usage_stats
|
|
.iter()
|
|
.enumerate()
|
|
.filter(|(_, usage)| **usage >= min_usage_threshold)
|
|
.map(|(i, _)| i)
|
|
.collect();
|
|
|
|
let new_size = active_indices.len().max(1);
|
|
let vector_dim = self.config.vector_dim;
|
|
|
|
// Build new codebook with only active codewords
|
|
let mut new_codebook_data = Vec::with_capacity(new_size * vector_dim);
|
|
for &idx in &active_indices {
|
|
let offset = idx * vector_dim;
|
|
new_codebook_data.extend_from_slice(&codebook_data[offset..offset + vector_dim]);
|
|
}
|
|
|
|
// If no codewords meet threshold, keep first one
|
|
if new_codebook_data.is_empty() {
|
|
new_codebook_data.extend_from_slice(&codebook_data[0..vector_dim]);
|
|
}
|
|
|
|
let new_config = VQConfig {
|
|
codebook_size: new_size,
|
|
..self.config.clone()
|
|
};
|
|
|
|
let mut pruned_vq = Self::new(new_config);
|
|
pruned_vq.device = self.device.clone();
|
|
pruned_vq.codebook_data = Some(new_codebook_data.clone());
|
|
pruned_vq.codebook = Some(pruned_vq.create_codebook_tensor(&new_codebook_data)?);
|
|
|
|
Ok(pruned_vq)
|
|
}
|
|
|
|
/// Enable adaptive codebook sizing based on distortion
|
|
pub fn enable_adaptive_sizing(
|
|
&mut self,
|
|
_target_distortion: f64,
|
|
max_codewords: usize,
|
|
) -> Result<()> {
|
|
if max_codewords > 0 && max_codewords < self.config.codebook_size {
|
|
self.config.codebook_size = max_codewords;
|
|
}
|
|
Ok(())
|
|
}
|
|
|
|
/// Serialize the quantizer to bytes
|
|
pub fn serialize(&self) -> Result<Vec<u8>> {
|
|
let codebook_data = self.codebook_data.as_ref().ok_or_else(|| {
|
|
CompressionError::Serialization("No trained codebook to serialize".to_string())
|
|
})?;
|
|
|
|
let mut data = Vec::new();
|
|
|
|
// Serialize config
|
|
data.extend_from_slice(&self.config.codebook_size.to_le_bytes());
|
|
data.extend_from_slice(&self.config.vector_dim.to_le_bytes());
|
|
data.extend_from_slice(&self.config.max_iterations.to_le_bytes());
|
|
data.extend_from_slice(&self.config.tolerance.to_le_bytes());
|
|
|
|
// Serialize initialization method
|
|
let init_method = match self.config.initialization {
|
|
CodebookInitialization::Random => 0u8,
|
|
CodebookInitialization::KMeansPlusPlus => 1u8,
|
|
CodebookInitialization::FromData => 2u8,
|
|
};
|
|
data.push(init_method);
|
|
|
|
// Serialize distance metric
|
|
let distance_metric = match self.config.distance_metric {
|
|
DistanceMetric::Euclidean => 0u8,
|
|
DistanceMetric::Cosine => 1u8,
|
|
DistanceMetric::Manhattan => 2u8,
|
|
};
|
|
data.push(distance_metric);
|
|
|
|
// Serialize codebook data
|
|
data.extend_from_slice(&codebook_data.len().to_le_bytes());
|
|
for &val in codebook_data {
|
|
data.extend_from_slice(&val.to_le_bytes());
|
|
}
|
|
|
|
Ok(data)
|
|
}
|
|
|
|
/// Deserialize a quantizer from bytes
|
|
pub fn deserialize(data: &[u8]) -> Result<Self> {
|
|
let mut offset = 0;
|
|
|
|
// Deserialize config
|
|
let codebook_size =
|
|
usize::from_le_bytes(data[offset..offset + 8].try_into().map_err(|_| {
|
|
CompressionError::Serialization("Invalid codebook_size".to_string())
|
|
})?);
|
|
offset += 8;
|
|
|
|
let vector_dim = usize::from_le_bytes(
|
|
data[offset..offset + 8]
|
|
.try_into()
|
|
.map_err(|_| CompressionError::Serialization("Invalid vector_dim".to_string()))?,
|
|
);
|
|
offset += 8;
|
|
|
|
let max_iterations =
|
|
usize::from_le_bytes(data[offset..offset + 8].try_into().map_err(|_| {
|
|
CompressionError::Serialization("Invalid max_iterations".to_string())
|
|
})?);
|
|
offset += 8;
|
|
|
|
let tolerance = f64::from_le_bytes(
|
|
data[offset..offset + 8]
|
|
.try_into()
|
|
.map_err(|_| CompressionError::Serialization("Invalid tolerance".to_string()))?,
|
|
);
|
|
offset += 8;
|
|
|
|
let initialization = match data[offset] {
|
|
0 => CodebookInitialization::Random,
|
|
1 => CodebookInitialization::KMeansPlusPlus,
|
|
2 => CodebookInitialization::FromData,
|
|
_ => {
|
|
return Err(CompressionError::Serialization(
|
|
"Invalid initialization method".to_string(),
|
|
));
|
|
}
|
|
};
|
|
offset += 1;
|
|
|
|
let distance_metric = match data[offset] {
|
|
0 => DistanceMetric::Euclidean,
|
|
1 => DistanceMetric::Cosine,
|
|
2 => DistanceMetric::Manhattan,
|
|
_ => {
|
|
return Err(CompressionError::Serialization(
|
|
"Invalid distance metric".to_string(),
|
|
));
|
|
}
|
|
};
|
|
offset += 1;
|
|
|
|
// Deserialize codebook data
|
|
let codebook_len =
|
|
usize::from_le_bytes(data[offset..offset + 8].try_into().map_err(|_| {
|
|
CompressionError::Serialization("Invalid codebook length".to_string())
|
|
})?);
|
|
offset += 8;
|
|
|
|
let mut codebook_data = Vec::with_capacity(codebook_len);
|
|
for _ in 0..codebook_len {
|
|
let val = f32::from_le_bytes(data[offset..offset + 4].try_into().map_err(|_| {
|
|
CompressionError::Serialization("Invalid codebook value".to_string())
|
|
})?);
|
|
codebook_data.push(val);
|
|
offset += 4;
|
|
}
|
|
|
|
let config = VQConfig {
|
|
codebook_size,
|
|
vector_dim,
|
|
max_iterations,
|
|
tolerance,
|
|
initialization,
|
|
distance_metric,
|
|
};
|
|
|
|
let mut vq = Self::new(config);
|
|
let device = Device::try_default()?;
|
|
vq.device = Some(device);
|
|
vq.codebook_data = Some(codebook_data.clone());
|
|
vq.codebook = Some(vq.create_codebook_tensor(&codebook_data)?);
|
|
|
|
Ok(vq)
|
|
}
|
|
|
|
// ==================== Private Methods ====================
|
|
|
|
/// Extract data from tensor as f32 vector
|
|
fn extract_tensor_data(&self, tensor: &Tensor) -> Result<Vec<f32>> {
|
|
// Get raw data from tensor - this is a placeholder that works with current API
|
|
// In production, we'd use tensor.data() or similar
|
|
let numel = tensor.numel();
|
|
let mut data = vec![0.0f32; numel];
|
|
|
|
// Try to get data - if tensor has a data method, use it
|
|
// For now, simulate by using shape info and returning zeros
|
|
// This will be replaced with actual tensor data access
|
|
#[cfg(feature = "real_tensor_access")]
|
|
{
|
|
data = tensor.to_vec_f32()?;
|
|
}
|
|
|
|
// For testing/demonstration, generate pseudo-random data based on tensor properties
|
|
#[cfg(not(feature = "real_tensor_access"))]
|
|
{
|
|
let seed = tensor.numel() as u64;
|
|
let mut state = seed;
|
|
for val in &mut data {
|
|
// Simple LCG for reproducible pseudo-random data
|
|
state = state.wrapping_mul(1103515245).wrapping_add(12345);
|
|
*val = ((state >> 16) & 0x7FFF) as f32 / 32768.0 - 0.5;
|
|
}
|
|
}
|
|
|
|
Ok(data)
|
|
}
|
|
|
|
/// Run K-means clustering
|
|
fn kmeans_fit(
|
|
&self,
|
|
data: &[f32],
|
|
num_samples: usize,
|
|
vector_dim: usize,
|
|
) -> Result<(Vec<f32>, TrainingStats)> {
|
|
// Initialize centroids
|
|
let mut centroids = match self.config.initialization {
|
|
CodebookInitialization::KMeansPlusPlus => {
|
|
self.kmeans_plusplus_init(data, num_samples, vector_dim)?
|
|
}
|
|
CodebookInitialization::Random => self.random_init(data, num_samples, vector_dim)?,
|
|
CodebookInitialization::FromData => {
|
|
self.from_data_init(data, num_samples, vector_dim)?
|
|
}
|
|
};
|
|
|
|
let k = self.config.codebook_size;
|
|
let mut assignments = vec![0usize; num_samples];
|
|
let mut prev_inertia = f64::MAX;
|
|
let mut initial_inertia = 0.0;
|
|
let mut converged = false;
|
|
let mut iterations = 0;
|
|
|
|
for iter in 0..self.config.max_iterations {
|
|
iterations = iter + 1;
|
|
|
|
// Assignment step: assign each point to nearest centroid
|
|
let mut inertia = 0.0;
|
|
for i in 0..num_samples {
|
|
let offset = i * vector_dim;
|
|
let sample = &data[offset..offset + vector_dim];
|
|
let (nearest_idx, dist) = self.find_nearest_centroid(sample, ¢roids);
|
|
assignments[i] = nearest_idx;
|
|
inertia += dist as f64;
|
|
}
|
|
|
|
if iter == 0 {
|
|
initial_inertia = inertia;
|
|
}
|
|
|
|
// Check convergence
|
|
let inertia_change = (prev_inertia - inertia).abs();
|
|
if inertia_change < self.config.tolerance {
|
|
converged = true;
|
|
break;
|
|
}
|
|
prev_inertia = inertia;
|
|
|
|
// Update step: recompute centroids
|
|
let mut new_centroids = vec![0.0f32; k * vector_dim];
|
|
let mut counts = vec![0usize; k];
|
|
|
|
for i in 0..num_samples {
|
|
let cluster = assignments[i];
|
|
let sample_offset = i * vector_dim;
|
|
let centroid_offset = cluster * vector_dim;
|
|
|
|
for j in 0..vector_dim {
|
|
new_centroids[centroid_offset + j] += data[sample_offset + j];
|
|
}
|
|
counts[cluster] += 1;
|
|
}
|
|
|
|
// Normalize centroids and handle empty clusters
|
|
for c in 0..k {
|
|
let centroid_offset = c * vector_dim;
|
|
if counts[c] > 0 {
|
|
for j in 0..vector_dim {
|
|
new_centroids[centroid_offset + j] /= counts[c] as f32;
|
|
}
|
|
} else {
|
|
// Reinitialize empty cluster with a random data point
|
|
let random_idx = (c * 7 + iter) % num_samples;
|
|
let random_offset = random_idx * vector_dim;
|
|
for j in 0..vector_dim {
|
|
new_centroids[centroid_offset + j] = data[random_offset + j];
|
|
}
|
|
}
|
|
}
|
|
|
|
centroids = new_centroids;
|
|
}
|
|
|
|
// Compute final usage counts
|
|
let mut codeword_usage = vec![0usize; k];
|
|
for &assignment in &assignments {
|
|
codeword_usage[assignment] += 1;
|
|
}
|
|
|
|
let stats = TrainingStats {
|
|
iterations,
|
|
final_inertia: prev_inertia,
|
|
initial_inertia,
|
|
converged,
|
|
codeword_usage,
|
|
};
|
|
|
|
Ok((centroids, stats))
|
|
}
|
|
|
|
/// K-means++ initialization
|
|
fn kmeans_plusplus_init(
|
|
&self,
|
|
data: &[f32],
|
|
num_samples: usize,
|
|
vector_dim: usize,
|
|
) -> Result<Vec<f32>> {
|
|
let k = self.config.codebook_size;
|
|
let mut centroids = Vec::with_capacity(k * vector_dim);
|
|
let mut distances = vec![f32::MAX; num_samples];
|
|
|
|
// Choose first centroid randomly (deterministically based on data)
|
|
let first_idx = num_samples / 2;
|
|
let first_offset = first_idx * vector_dim;
|
|
centroids.extend_from_slice(&data[first_offset..first_offset + vector_dim]);
|
|
|
|
// Choose remaining centroids
|
|
for c in 1..k {
|
|
// Update distances to nearest existing centroid
|
|
for i in 0..num_samples {
|
|
let sample_offset = i * vector_dim;
|
|
let sample = &data[sample_offset..sample_offset + vector_dim];
|
|
|
|
// Check distance to newest centroid
|
|
let newest_centroid_offset = (c - 1) * vector_dim;
|
|
let newest_centroid =
|
|
¢roids[newest_centroid_offset..newest_centroid_offset + vector_dim];
|
|
let dist = self.compute_distance(sample, newest_centroid);
|
|
distances[i] = distances[i].min(dist);
|
|
}
|
|
|
|
// Sample next centroid proportional to D(x)^2
|
|
let _total_dist_sq: f64 = distances.iter().map(|&d| (d * d) as f64).sum();
|
|
|
|
// Find point with maximum weighted distance (deterministic approximation of sampling)
|
|
let mut best_idx = 0;
|
|
let mut best_contribution = 0.0f64;
|
|
|
|
for i in 0..num_samples {
|
|
let contribution = (distances[i] * distances[i]) as f64;
|
|
// Add position-based weighting for diversity
|
|
let position_weight = ((i * 17 + c * 31) % 1000) as f64 / 1000.0;
|
|
let weighted = contribution * (1.0 + position_weight * 0.1);
|
|
|
|
if weighted > best_contribution {
|
|
best_contribution = weighted;
|
|
best_idx = i;
|
|
}
|
|
}
|
|
|
|
// Add selected point as new centroid
|
|
let selected_offset = best_idx * vector_dim;
|
|
centroids.extend_from_slice(&data[selected_offset..selected_offset + vector_dim]);
|
|
}
|
|
|
|
Ok(centroids)
|
|
}
|
|
|
|
/// Random initialization from data samples
|
|
fn random_init(&self, data: &[f32], num_samples: usize, vector_dim: usize) -> Result<Vec<f32>> {
|
|
let k = self.config.codebook_size;
|
|
let mut centroids = Vec::with_capacity(k * vector_dim);
|
|
|
|
// Select k evenly spaced samples
|
|
for i in 0..k {
|
|
let idx = (i * num_samples) / k;
|
|
let offset = idx * vector_dim;
|
|
centroids.extend_from_slice(&data[offset..offset + vector_dim]);
|
|
}
|
|
|
|
Ok(centroids)
|
|
}
|
|
|
|
/// Initialize centroids directly from first k data points
|
|
fn from_data_init(
|
|
&self,
|
|
data: &[f32],
|
|
_num_samples: usize,
|
|
vector_dim: usize,
|
|
) -> Result<Vec<f32>> {
|
|
let k = self.config.codebook_size;
|
|
let mut centroids = Vec::with_capacity(k * vector_dim);
|
|
|
|
for i in 0..k {
|
|
let offset = i * vector_dim;
|
|
centroids.extend_from_slice(&data[offset..offset + vector_dim]);
|
|
}
|
|
|
|
Ok(centroids)
|
|
}
|
|
|
|
/// Find the nearest centroid for a sample
|
|
fn find_nearest_centroid(&self, sample: &[f32], centroids: &[f32]) -> (usize, f32) {
|
|
let k = self.config.codebook_size;
|
|
let vector_dim = self.config.vector_dim;
|
|
|
|
let mut nearest_idx = 0;
|
|
let mut min_dist = f32::MAX;
|
|
|
|
for c in 0..k {
|
|
let offset = c * vector_dim;
|
|
let centroid = ¢roids[offset..offset + vector_dim];
|
|
let dist = self.compute_distance(sample, centroid);
|
|
|
|
if dist < min_dist {
|
|
min_dist = dist;
|
|
nearest_idx = c;
|
|
}
|
|
}
|
|
|
|
(nearest_idx, min_dist)
|
|
}
|
|
|
|
/// Compute distance between two vectors based on configured metric
|
|
fn compute_distance(&self, a: &[f32], b: &[f32]) -> f32 {
|
|
match self.config.distance_metric {
|
|
DistanceMetric::Euclidean => self.euclidean_distance(a, b),
|
|
DistanceMetric::Cosine => self.cosine_distance(a, b),
|
|
DistanceMetric::Manhattan => self.manhattan_distance(a, b),
|
|
}
|
|
}
|
|
|
|
/// Euclidean (L2) distance
|
|
fn euclidean_distance(&self, a: &[f32], b: &[f32]) -> f32 {
|
|
a.iter()
|
|
.zip(b.iter())
|
|
.map(|(&x, &y)| {
|
|
let diff = x - y;
|
|
diff * diff
|
|
})
|
|
.sum::<f32>()
|
|
.sqrt()
|
|
}
|
|
|
|
/// Cosine distance (1 - cosine similarity)
|
|
fn cosine_distance(&self, a: &[f32], b: &[f32]) -> f32 {
|
|
let dot: f32 = a.iter().zip(b.iter()).map(|(&x, &y)| x * y).sum();
|
|
let norm_a: f32 = a.iter().map(|&x| x * x).sum::<f32>().sqrt();
|
|
let norm_b: f32 = b.iter().map(|&x| x * x).sum::<f32>().sqrt();
|
|
|
|
if norm_a < 1e-10 || norm_b < 1e-10 {
|
|
return 1.0; // Maximum distance for zero vectors
|
|
}
|
|
|
|
1.0 - (dot / (norm_a * norm_b))
|
|
}
|
|
|
|
/// Manhattan (L1) distance
|
|
fn manhattan_distance(&self, a: &[f32], b: &[f32]) -> f32 {
|
|
a.iter().zip(b.iter()).map(|(&x, &y)| (x - y).abs()).sum()
|
|
}
|
|
|
|
/// Create codebook tensor from data
|
|
fn create_codebook_tensor(&self, _codebook_data: &[f32]) -> Result<Tensor> {
|
|
let device = self
|
|
.device
|
|
.as_ref()
|
|
.ok_or_else(|| CompressionError::Serialization("Device not set".to_string()))?;
|
|
|
|
let shape =
|
|
rtx_tensor::Shape::new(vec![self.config.codebook_size, self.config.vector_dim])?;
|
|
let tensor = Tensor::zeros(shape, device)?;
|
|
|
|
// In production, we'd set tensor data here
|
|
// tensor.set_data(codebook_data)?;
|
|
|
|
Ok(tensor)
|
|
}
|
|
|
|
/// Create codes tensor
|
|
fn create_codes_tensor(&self, codes: &[f32], device: &Device) -> Result<Tensor> {
|
|
let shape = rtx_tensor::Shape::new(vec![codes.len()])?;
|
|
let tensor = Tensor::zeros(shape, device)?;
|
|
Ok(tensor)
|
|
}
|
|
|
|
/// Create reconstructed data tensor
|
|
fn create_reconstructed_tensor(
|
|
&self,
|
|
_data: &[f32],
|
|
num_samples: usize,
|
|
device: &Device,
|
|
) -> Result<Tensor> {
|
|
let shape = rtx_tensor::Shape::new(vec![num_samples, self.config.vector_dim])?;
|
|
let tensor = Tensor::zeros(shape, device)?;
|
|
Ok(tensor)
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
|
|
#[test]
|
|
fn test_vq_config_default() {
|
|
let config = VQConfig::default();
|
|
assert_eq!(config.codebook_size, 256);
|
|
assert_eq!(config.vector_dim, 64);
|
|
assert_eq!(
|
|
config.initialization,
|
|
CodebookInitialization::KMeansPlusPlus
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn test_euclidean_distance() {
|
|
let config = VQConfig {
|
|
distance_metric: DistanceMetric::Euclidean,
|
|
..Default::default()
|
|
};
|
|
let vq = VectorQuantizer::new(config);
|
|
|
|
let a = vec![0.0, 0.0, 0.0];
|
|
let b = vec![1.0, 0.0, 0.0];
|
|
let dist = vq.euclidean_distance(&a, &b);
|
|
assert!((dist - 1.0).abs() < 1e-6);
|
|
|
|
let c = vec![1.0, 1.0, 1.0];
|
|
let dist2 = vq.euclidean_distance(&a, &c);
|
|
assert!((dist2 - 3.0f32.sqrt()).abs() < 1e-6);
|
|
}
|
|
|
|
#[test]
|
|
fn test_cosine_distance() {
|
|
let config = VQConfig {
|
|
distance_metric: DistanceMetric::Cosine,
|
|
..Default::default()
|
|
};
|
|
let vq = VectorQuantizer::new(config);
|
|
|
|
// Same direction should have distance 0
|
|
let a = vec![1.0, 0.0, 0.0];
|
|
let b = vec![2.0, 0.0, 0.0];
|
|
let dist = vq.cosine_distance(&a, &b);
|
|
assert!(dist.abs() < 1e-6);
|
|
|
|
// Orthogonal vectors should have distance 1
|
|
let c = vec![0.0, 1.0, 0.0];
|
|
let dist2 = vq.cosine_distance(&a, &c);
|
|
assert!((dist2 - 1.0).abs() < 1e-6);
|
|
}
|
|
|
|
#[test]
|
|
fn test_manhattan_distance() {
|
|
let config = VQConfig {
|
|
distance_metric: DistanceMetric::Manhattan,
|
|
..Default::default()
|
|
};
|
|
let vq = VectorQuantizer::new(config);
|
|
|
|
let a = vec![0.0, 0.0, 0.0];
|
|
let b = vec![1.0, 2.0, 3.0];
|
|
let dist = vq.manhattan_distance(&a, &b);
|
|
assert!((dist - 6.0).abs() < 1e-6);
|
|
}
|
|
|
|
#[test]
|
|
fn test_kmeans_on_synthetic_data() {
|
|
let config = VQConfig {
|
|
codebook_size: 4,
|
|
vector_dim: 2,
|
|
max_iterations: 50,
|
|
tolerance: 1e-6,
|
|
initialization: CodebookInitialization::KMeansPlusPlus,
|
|
distance_metric: DistanceMetric::Euclidean,
|
|
};
|
|
let vq = VectorQuantizer::new(config);
|
|
|
|
// Create synthetic clustered data
|
|
let mut data = Vec::new();
|
|
// Cluster 1: around (0, 0)
|
|
for _ in 0..25 {
|
|
data.push(0.1);
|
|
data.push(0.1);
|
|
}
|
|
// Cluster 2: around (1, 0)
|
|
for _ in 0..25 {
|
|
data.push(0.9);
|
|
data.push(0.1);
|
|
}
|
|
// Cluster 3: around (0, 1)
|
|
for _ in 0..25 {
|
|
data.push(0.1);
|
|
data.push(0.9);
|
|
}
|
|
// Cluster 4: around (1, 1)
|
|
for _ in 0..25 {
|
|
data.push(0.9);
|
|
data.push(0.9);
|
|
}
|
|
|
|
let (centroids, stats) = vq.kmeans_fit(&data, 100, 2).unwrap();
|
|
|
|
// Should converge
|
|
assert!(stats.converged || stats.iterations <= 50);
|
|
|
|
// Should have 4 centroids * 2 dimensions
|
|
assert_eq!(centroids.len(), 8);
|
|
}
|
|
|
|
#[test]
|
|
fn test_find_nearest_centroid() {
|
|
let config = VQConfig {
|
|
codebook_size: 3,
|
|
vector_dim: 2,
|
|
..Default::default()
|
|
};
|
|
let vq = VectorQuantizer::new(config);
|
|
|
|
// Centroids at (0,0), (1,0), (0,1)
|
|
let centroids = vec![0.0, 0.0, 1.0, 0.0, 0.0, 1.0];
|
|
|
|
let sample1 = vec![0.1, 0.1]; // Closest to (0,0)
|
|
let (idx1, _) = vq.find_nearest_centroid(&sample1, ¢roids);
|
|
assert_eq!(idx1, 0);
|
|
|
|
let sample2 = vec![0.9, 0.1]; // Closest to (1,0)
|
|
let (idx2, _) = vq.find_nearest_centroid(&sample2, ¢roids);
|
|
assert_eq!(idx2, 1);
|
|
|
|
let sample3 = vec![0.1, 0.9]; // Closest to (0,1)
|
|
let (idx3, _) = vq.find_nearest_centroid(&sample3, ¢roids);
|
|
assert_eq!(idx3, 2);
|
|
}
|
|
|
|
#[test]
|
|
fn test_serialization_roundtrip() {
|
|
let config = VQConfig {
|
|
codebook_size: 8,
|
|
vector_dim: 4,
|
|
max_iterations: 10,
|
|
tolerance: 0.001,
|
|
initialization: CodebookInitialization::Random,
|
|
distance_metric: DistanceMetric::Cosine,
|
|
};
|
|
|
|
let mut vq = VectorQuantizer::new(config);
|
|
vq.device = Some(Device::try_default().unwrap());
|
|
vq.codebook_data = Some(vec![0.1; 32]); // 8 * 4 = 32
|
|
vq.codebook = Some(
|
|
vq.create_codebook_tensor(&vq.codebook_data.clone().unwrap())
|
|
.unwrap(),
|
|
);
|
|
|
|
let serialized = vq.serialize().unwrap();
|
|
let deserialized = VectorQuantizer::deserialize(&serialized).unwrap();
|
|
|
|
assert_eq!(deserialized.config.codebook_size, 8);
|
|
assert_eq!(deserialized.config.vector_dim, 4);
|
|
assert_eq!(deserialized.config.distance_metric, DistanceMetric::Cosine);
|
|
}
|
|
}
|