//! 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, } /// Vector quantizer using K-means clustering pub struct VectorQuantizer { config: VQConfig, /// Codebook stored as flattened vector [codebook_size * vector_dim] codebook_data: Option>, codebook: Option, device: Option, online_mode: bool, learning_rate: f64, training_stats: Option, } 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 { 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 { 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> { 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 = 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 { 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::>() } else { vec![1.0 / self.config.codebook_size as f64; self.config.codebook_size] }; // Select codewords that meet threshold let active_indices: Vec = 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> { 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 { 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> { // 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, 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> { 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> { 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> { 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::() .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::().sqrt(); let norm_b: f32 = b.iter().map(|&x| x * x).sum::().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 { 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 { 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 { 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); } }