Skip to main content

ohms_adaptq/novaq/
codebooks.rs

1use crate::Result;
2use super::{WeightMatrix, VectorCodebooks, QuantizationIndices, CodebookEntry, SubspaceStrategy, FallbackQuantizer, QuantizationStrategy};
3use rand::{Rng, SeedableRng};
4use rand_chacha::ChaCha8Rng;
5
6/// Multi-stage Vector Codebook Builder implementing Stage 2 of NOVAQ
7/// 
8/// Mathematical formulation:
9/// For each subspace k:
10/// 1. Coarse quantization: b^(1)_{i,k} = argmin_c ||v_{i,k} - C^(1)_{c,k}||²
11/// 2. Residual calculation: r_{i,k} = v_{i,k} - C^(1)_{b^(1)_{i,k},k}
12/// 3. Residual quantization: b^(2)_{i,k} = argmin_c ||r_{i,k} - C^(2)_{c,k}||²
13/// 
14/// Final reconstruction: W_{i,:} = Σ_{k=1}^N (C^(1)_{b^(1)_{i,k},k} + C^(2)_{b^(2)_{i,k},k})
15#[derive(Debug)]
16pub struct CodebookBuilder {
17    num_subspaces: usize,
18    codebook_size_l1: usize,
19    codebook_size_l2: usize,
20    rng: ChaCha8Rng,
21    max_kmeans_iterations: usize,
22    convergence_threshold: f32,
23    subspace_strategy: SubspaceStrategy,
24    fallback_quantizer: FallbackQuantizer,
25}
26
27impl CodebookBuilder {
28    pub fn new(num_subspaces: usize, codebook_size_l1: usize, codebook_size_l2: usize, seed: u64) -> Self {
29        Self {
30            num_subspaces,
31            codebook_size_l1,
32            codebook_size_l2,
33            rng: ChaCha8Rng::seed_from_u64(seed),
34            max_kmeans_iterations: 100,
35            convergence_threshold: 1e-6,
36            subspace_strategy: SubspaceStrategy::new(num_subspaces, codebook_size_l1, codebook_size_l2),
37            fallback_quantizer: FallbackQuantizer::new(),
38        }
39    }
40    
41    /// Build multi-stage vector codebooks for given weight matrix with enhanced stability
42    pub fn build_codebooks(&mut self, weights: &WeightMatrix) -> Result<(VectorCodebooks, QuantizationIndices)> {
43        let rows = weights.rows();
44        let cols = weights.cols();
45        
46        // Use enhanced subspace strategy to determine optimal configuration
47        let config = self.subspace_strategy.determine_config(rows, cols);
48        
49        // Print configuration for user feedback
50        self.subspace_strategy.print_config(&config, rows, cols);
51        
52        // Handle edge cases with fallback quantization
53        if config.strategy == QuantizationStrategy::ScalarQuantization {
54            return self.fallback_quantizer.scalar_quantize(weights, &config);
55        }
56        
57        // Use the determined configuration for quantization
58        let effective_subspaces = config.effective_subspaces;
59        let subspace_size = config.subspace_size;
60        let l1_size = config.codebook_size_l1;
61        let l2_size = config.codebook_size_l2;
62        
63        let mut level1_codebooks = Vec::with_capacity(effective_subspaces);
64        let mut level2_codebooks = Vec::with_capacity(effective_subspaces);
65        let mut level1_indices = vec![vec![0u8; effective_subspaces]; rows];
66        let mut level2_indices = vec![vec![0u8; effective_subspaces]; rows];
67        
68        // Process each subspace with enhanced error handling
69        for subspace_idx in 0..effective_subspaces {
70            let start_col = subspace_idx * subspace_size;
71            let end_col = start_col + subspace_size;
72            
73            // Bounds checking to prevent array access errors
74            if end_col > cols {
75                return Err(format!("Subspace bounds error: end_col {} > cols {}", end_col, cols).into());
76            }
77            
78            // Extract subspace vectors
79            let mut subspace_vectors = Vec::with_capacity(rows);
80            for row_idx in 0..rows {
81                let row = weights.get_row(row_idx);
82                if end_col <= row.len() {
83                    subspace_vectors.push(row[start_col..end_col].to_vec());
84                } else {
85                    return Err(format!("Row bounds error: end_col {} > row.len() {}", end_col, row.len()).into());
86                }
87            }
88            
89            // Stage 1: Build first-level codebook with enhanced parameters
90            let l1_codebook = self.build_kmeans_codebook(&subspace_vectors, l1_size)?;
91            
92            // Stage 2: Calculate residuals and build second-level codebook
93            let mut residuals = Vec::with_capacity(rows);
94            for (row_idx, vector) in subspace_vectors.iter().enumerate() {
95                // Find nearest centroid in L1 codebook
96                let (nearest_idx, _) = self.find_nearest_centroid(vector, &l1_codebook);
97                level1_indices[row_idx][subspace_idx] = nearest_idx as u8;
98                
99                // Calculate residual: r_{i,k} = v_{i,k} - C^(1)_{b^(1)_{i,k},k}
100                let residual = self.calculate_residual(vector, &l1_codebook[nearest_idx].centroid);
101                residuals.push(residual);
102            }
103            
104            // Build second-level codebook for residuals with enhanced parameters
105            let l2_codebook = self.build_kmeans_codebook(&residuals, l2_size)?;
106            
107            // Assign residuals to L2 codebook
108            for (row_idx, residual) in residuals.iter().enumerate() {
109                let (nearest_idx, _) = self.find_nearest_centroid(residual, &l2_codebook);
110                level2_indices[row_idx][subspace_idx] = nearest_idx as u8;
111            }
112            
113            level1_codebooks.push(l1_codebook);
114            level2_codebooks.push(l2_codebook);
115        }
116        
117        let codebooks = VectorCodebooks {
118            level1_codebooks,
119            level2_codebooks,
120            subspace_size,
121        };
122        
123        let indices = QuantizationIndices {
124            level1_indices,
125            level2_indices,
126        };
127        
128        println!("✅ Codebook generation completed successfully");
129        Ok((codebooks, indices))
130    }
131    
132    /// Reconstruct weights from codebooks and indices
133    pub fn reconstruct_weights(
134        &self,
135        codebooks: &VectorCodebooks,
136        indices: &QuantizationIndices,
137        rows: usize,
138        cols: usize,
139    ) -> Result<Vec<f32>> {
140        let mut reconstructed = vec![0.0; rows * cols];
141        
142        for row_idx in 0..rows {
143            for subspace_idx in 0..self.num_subspaces {
144                let start_col = subspace_idx * codebooks.subspace_size;
145                
146                // Get codebook entries
147                let l1_idx = indices.level1_indices[row_idx][subspace_idx] as usize;
148                let l2_idx = indices.level2_indices[row_idx][subspace_idx] as usize;
149                
150                let l1_centroid = &codebooks.level1_codebooks[subspace_idx][l1_idx].centroid;
151                let l2_centroid = &codebooks.level2_codebooks[subspace_idx][l2_idx].centroid;
152                
153                // Reconstruct: v = C^(1) + C^(2)
154                for (offset, (&l1_val, &l2_val)) in l1_centroid.iter().zip(l2_centroid.iter()).enumerate() {
155                    let col_idx = start_col + offset;
156                    let flat_idx = row_idx * cols + col_idx;
157                    reconstructed[flat_idx] = l1_val + l2_val;
158                }
159            }
160        }
161        
162        Ok(reconstructed)
163    }
164    
165    /// Build codebook using k-means clustering
166    fn build_kmeans_codebook(&mut self, vectors: &[Vec<f32>], k: usize) -> Result<Vec<CodebookEntry>> {
167        if vectors.is_empty() {
168            return Err("Cannot build codebook from empty vector set".into());
169        }
170        
171        let vector_dim = vectors[0].len();
172        let n_vectors = vectors.len();
173        
174        // Handle case where k >= number of vectors
175        let effective_k = k.min(n_vectors);
176        
177        // Initialize centroids randomly
178        let mut centroids = Vec::with_capacity(effective_k);
179        for _ in 0..effective_k {
180            let mut centroid = vec![0.0; vector_dim];
181            for dim in 0..vector_dim {
182                centroid[dim] = self.rng.gen_range(-1.0..1.0);
183            }
184            centroids.push(centroid);
185        }
186        
187        let mut assignments = vec![0; n_vectors];
188        let mut prev_centroids = centroids.clone();
189        
190        // K-means iterations
191        for iteration in 0..self.max_kmeans_iterations {
192            // Assignment step: assign each vector to nearest centroid
193            for (vec_idx, vector) in vectors.iter().enumerate() {
194                let (nearest_idx, _) = self.find_nearest_centroid_raw(vector, &centroids);
195                assignments[vec_idx] = nearest_idx;
196            }
197            
198            // Update step: recompute centroids
199            for cluster_idx in 0..effective_k {
200                let cluster_vectors: Vec<&Vec<f32>> = vectors.iter()
201                    .enumerate()
202                    .filter(|(idx, _)| assignments[*idx] == cluster_idx)
203                    .map(|(_, vec)| vec)
204                    .collect();
205                
206                if !cluster_vectors.is_empty() {
207                    // Compute mean of assigned vectors
208                    for dim in 0..vector_dim {
209                        let sum: f32 = cluster_vectors.iter().map(|v| v[dim]).sum();
210                        centroids[cluster_idx][dim] = sum / cluster_vectors.len() as f32;
211                    }
212                } else {
213                    // Reinitialize empty cluster
214                    for dim in 0..vector_dim {
215                        centroids[cluster_idx][dim] = self.rng.gen_range(-1.0..1.0);
216                    }
217                }
218            }
219            
220            // Check convergence
221            let converged = self.check_convergence(&prev_centroids, &centroids);
222            if converged {
223                break;
224            }
225            
226            prev_centroids = centroids.clone();
227        }
228        
229        // Create codebook entries with usage counts
230        let mut codebook = Vec::with_capacity(effective_k);
231        for cluster_idx in 0..effective_k {
232            let usage_count = assignments.iter().filter(|&&idx| idx == cluster_idx).count();
233            codebook.push(CodebookEntry {
234                centroid: centroids[cluster_idx].clone(),
235                usage_count,
236            });
237        }
238        
239        Ok(codebook)
240    }
241    
242    /// Find nearest centroid and return index and distance
243    fn find_nearest_centroid(&self, vector: &[f32], codebook: &[CodebookEntry]) -> (usize, f32) {
244        let centroids: Vec<Vec<f32>> = codebook.iter().map(|entry| entry.centroid.clone()).collect();
245        self.find_nearest_centroid_raw(vector, &centroids)
246    }
247    
248    /// Find nearest centroid in raw centroid list
249    fn find_nearest_centroid_raw(&self, vector: &[f32], centroids: &[Vec<f32>]) -> (usize, f32) {
250        let mut min_distance = f32::INFINITY;
251        let mut nearest_idx = 0;
252        
253        for (idx, centroid) in centroids.iter().enumerate() {
254            let distance = self.euclidean_distance(vector, centroid);
255            if distance < min_distance {
256                min_distance = distance;
257                nearest_idx = idx;
258            }
259        }
260        
261        (nearest_idx, min_distance)
262    }
263    
264    /// Calculate residual vector
265    fn calculate_residual(&self, vector: &[f32], centroid: &[f32]) -> Vec<f32> {
266        vector.iter().zip(centroid.iter()).map(|(&v, &c)| v - c).collect()
267    }
268    
269    /// Calculate Euclidean distance between two vectors
270    fn euclidean_distance(&self, a: &[f32], b: &[f32]) -> f32 {
271        a.iter().zip(b.iter()).map(|(&x, &y)| (x - y).powi(2)).sum::<f32>().sqrt()
272    }
273    
274    /// Check convergence of k-means algorithm
275    fn check_convergence(&self, prev_centroids: &[Vec<f32>], current_centroids: &[Vec<f32>]) -> bool {
276        for (prev, current) in prev_centroids.iter().zip(current_centroids.iter()) {
277            let distance = self.euclidean_distance(prev, current);
278            if distance > self.convergence_threshold {
279                return false;
280            }
281        }
282        true
283    }
284}
285
286/// Calculate quantization error metrics
287pub struct QuantizationMetrics {
288    pub mse: f32,
289    pub psnr: f32,
290    pub max_error: f32,
291    pub mean_error: f32,
292}
293
294impl QuantizationMetrics {
295    pub fn calculate(original: &[f32], reconstructed: &[f32]) -> Self {
296        assert_eq!(original.len(), reconstructed.len());
297        
298        let n = original.len() as f32;
299        let mut mse = 0.0;
300        let mut max_error: f32 = 0.0;
301        let mut total_error = 0.0;
302        
303        for (&orig, &recon) in original.iter().zip(reconstructed.iter()) {
304            let error = (orig - recon).abs();
305            let squared_error = error * error;
306            
307            mse += squared_error;
308            max_error = max_error.max(error);
309            total_error += error;
310        }
311        
312        mse /= n;
313        let mean_error = total_error / n;
314        
315        // Calculate PSNR (Peak Signal-to-Noise Ratio)
316        let max_val = original.iter().fold(0.0f32, |acc, &x| acc.max(x.abs()));
317        let psnr = if mse > 0.0 {
318            20.0 * (max_val / mse.sqrt()).log10()
319        } else {
320            f32::INFINITY
321        };
322        
323        Self {
324            mse,
325            psnr,
326            max_error,
327            mean_error,
328        }
329    }
330    
331    pub fn print_summary(&self) {
332        println!("Quantization Metrics:");
333        println!("  MSE: {:.6}", self.mse);
334        println!("  PSNR: {:.2} dB", self.psnr);
335        println!("  Max error: {:.6}", self.max_error);
336        println!("  Mean error: {:.6}", self.mean_error);
337    }
338}
339
340#[cfg(test)]
341mod tests {
342    use super::*;
343    
344    #[test]
345    fn test_codebook_builder_creation() {
346        let builder = CodebookBuilder::new(4, 16, 4, 42);
347        assert_eq!(builder.num_subspaces, 4);
348        assert_eq!(builder.codebook_size_l1, 16);
349        assert_eq!(builder.codebook_size_l2, 4);
350    }
351    
352    #[test]
353    fn test_build_codebooks_basic() {
354        let data = vec![
355            1.0, 2.0, 3.0, 4.0,  // 4 columns for 2 subspaces of size 2
356            5.0, 6.0, 7.0, 8.0,
357            1.1, 2.1, 3.1, 4.1,
358            5.1, 6.1, 7.1, 8.1,
359        ];
360        let weights = WeightMatrix::new(data, vec![4, 4], "test".to_string());
361        
362        let mut builder = CodebookBuilder::new(2, 4, 2, 42);
363        let result = builder.build_codebooks(&weights);
364        
365        assert!(result.is_ok());
366        let (codebooks, indices) = result.unwrap();
367        
368        assert_eq!(codebooks.level1_codebooks.len(), 2);
369        assert_eq!(codebooks.level2_codebooks.len(), 2);
370        assert_eq!(codebooks.subspace_size, 2);
371        assert_eq!(indices.level1_indices.len(), 4);
372        assert_eq!(indices.level2_indices.len(), 4);
373    }
374    
375    #[test]
376    fn test_reconstruction() {
377        let data = vec![
378            1.0, 2.0, 3.0, 4.0,
379            5.0, 6.0, 7.0, 8.0,
380        ];
381        let weights = WeightMatrix::new(data.clone(), vec![2, 4], "test".to_string());
382        
383        let mut builder = CodebookBuilder::new(2, 2, 2, 42);
384        let (codebooks, indices) = builder.build_codebooks(&weights).unwrap();
385        
386        let reconstructed = builder.reconstruct_weights(&codebooks, &indices, 2, 4).unwrap();
387        
388        assert_eq!(reconstructed.len(), 8);
389        
390        // Calculate reconstruction error
391        let metrics = QuantizationMetrics::calculate(&data, &reconstructed);
392        println!("Reconstruction MSE: {:.6}", metrics.mse);
393        
394        // Should have reasonable reconstruction quality
395        assert!(metrics.mse < 10.0); // Adjust threshold as needed
396    }
397    
398    #[test]
399    fn test_euclidean_distance() {
400        let builder = CodebookBuilder::new(2, 2, 2, 42);
401        let a = vec![1.0, 2.0, 3.0];
402        let b = vec![4.0, 5.0, 6.0];
403        
404        let distance = builder.euclidean_distance(&a, &b);
405        let expected = ((3.0_f32).powi(2) + (3.0_f32).powi(2) + (3.0_f32).powi(2)).sqrt();
406        
407        assert!((distance - expected).abs() < 1e-6);
408    }
409    
410    #[test]
411    fn test_residual_calculation() {
412        let builder = CodebookBuilder::new(2, 2, 2, 42);
413        let vector = vec![5.0, 7.0, 9.0];
414        let centroid = vec![2.0, 3.0, 4.0];
415        
416        let residual = builder.calculate_residual(&vector, &centroid);
417        let expected = vec![3.0, 4.0, 5.0];
418        
419        assert_eq!(residual, expected);
420    }
421}