Initial commit
This commit is contained in:
@@ -0,0 +1,997 @@
|
||||
//! 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);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user