Skip to main content

ohms_adaptq/novaq/
subspace_strategy.rs

1use crate::Result;
2use super::{WeightMatrix, VectorCodebooks, QuantizationIndices, CodebookEntry, NumericalStabilityGuard};
3
4/// Fallback quantization strategies for edge cases
5#[derive(Debug, Clone, Copy, PartialEq)]
6pub enum QuantizationStrategy {
7    /// Standard multi-subspace quantization
8    MultiSubspace,
9    /// Single-subspace with enhanced codebook size
10    SingleSubspaceEnhanced,
11    /// Direct scalar quantization for very small tensors
12    ScalarQuantization,
13    /// Uniform quantization with dithering
14    UniformDithered,
15}
16
17/// Enhanced subspace configuration that handles edge cases robustly
18#[derive(Debug, Clone)]
19pub struct SubspaceConfig {
20    pub effective_subspaces: usize,
21    pub subspace_size: usize,
22    pub strategy: QuantizationStrategy,
23    pub codebook_size_l1: usize,
24    pub codebook_size_l2: usize,
25    pub min_vectors_per_cluster: usize,
26}
27
28/// Robust subspace strategy selector that prevents mathematical instabilities
29#[derive(Debug)]
30pub struct SubspaceStrategy {
31    pub original_subspaces: usize,
32    pub original_codebook_l1: usize,
33    pub original_codebook_l2: usize,
34    pub min_subspace_size: usize,
35    pub min_vectors_per_cluster: usize,
36}
37
38impl SubspaceStrategy {
39    pub fn new(
40        original_subspaces: usize,
41        original_codebook_l1: usize,
42        original_codebook_l2: usize,
43    ) -> Self {
44        Self {
45            original_subspaces,
46            original_codebook_l1,
47            original_codebook_l2,
48            min_subspace_size: 2, // Minimum for mathematical stability
49            min_vectors_per_cluster: 3, // Minimum for meaningful clustering
50        }
51    }
52
53    /// Determine optimal subspace configuration based on tensor dimensions and constraints
54    pub fn determine_config(
55        &self,
56        rows: usize,
57        cols: usize,
58    ) -> SubspaceConfig {
59        // Handle degenerate cases first
60        if rows == 0 || cols == 0 {
61            return SubspaceConfig {
62                effective_subspaces: 1,
63                subspace_size: 1,
64                strategy: QuantizationStrategy::ScalarQuantization,
65                codebook_size_l1: 2,
66                codebook_size_l2: 2,
67                min_vectors_per_cluster: 1,
68            };
69        }
70
71        // For very small tensors, use scalar quantization
72        if cols == 1 || rows < self.min_vectors_per_cluster {
73            return self.create_scalar_config(rows, cols);
74        }
75
76        // For small tensors that can't support multiple subspaces
77        if cols < self.original_subspaces * self.min_subspace_size {
78            return self.create_single_subspace_config(rows, cols);
79        }
80
81        // Find optimal subspace division for normal cases
82        self.create_multi_subspace_config(rows, cols)
83    }
84
85    /// Create configuration for scalar quantization (very small tensors)
86    fn create_scalar_config(&self, rows: usize, cols: usize) -> SubspaceConfig {
87        let codebook_size = if rows >= 8 {
88            self.original_codebook_l1.min(16)
89        } else {
90            (rows / 2).max(2).min(8)
91        };
92
93        SubspaceConfig {
94            effective_subspaces: 1,
95            subspace_size: cols,
96            strategy: QuantizationStrategy::ScalarQuantization,
97            codebook_size_l1: codebook_size,
98            codebook_size_l2: 2, // Minimal residual codebook
99            min_vectors_per_cluster: (rows / codebook_size).max(1),
100        }
101    }
102
103    /// Create configuration for single subspace with enhanced codebook
104    fn create_single_subspace_config(&self, rows: usize, cols: usize) -> SubspaceConfig {
105        // Use larger codebook to compensate for lack of subspaces
106        let enhanced_l1_size = if cols >= 4 {
107            (self.original_codebook_l1 * 2).min(64)
108        } else {
109            self.original_codebook_l1
110        };
111
112        let enhanced_l2_size = if cols >= 4 {
113            (self.original_codebook_l2 * 2).min(16)
114        } else {
115            self.original_codebook_l2
116        };
117
118        SubspaceConfig {
119            effective_subspaces: 1,
120            subspace_size: cols,
121            strategy: QuantizationStrategy::SingleSubspaceEnhanced,
122            codebook_size_l1: enhanced_l1_size,
123            codebook_size_l2: enhanced_l2_size,
124            min_vectors_per_cluster: (rows / enhanced_l1_size).max(1),
125        }
126    }
127
128    /// Create configuration for multi-subspace quantization (normal case)
129    fn create_multi_subspace_config(&self, rows: usize, cols: usize) -> SubspaceConfig {
130        let mut best_config = SubspaceConfig {
131            effective_subspaces: 1,
132            subspace_size: cols,
133            strategy: QuantizationStrategy::SingleSubspaceEnhanced,
134            codebook_size_l1: self.original_codebook_l1,
135            codebook_size_l2: self.original_codebook_l2,
136            min_vectors_per_cluster: rows / self.original_codebook_l1,
137        };
138
139        // Try different subspace counts to find the best fit
140        for num_subspaces in (2..=self.original_subspaces).rev() {
141            if cols % num_subspaces == 0 {
142                let subspace_size = cols / num_subspaces;
143                
144                // Ensure subspace size meets minimum requirements
145                if subspace_size >= self.min_subspace_size {
146                    // Calculate minimum vectors needed per cluster
147                    let vectors_per_cluster = rows / self.original_codebook_l1;
148                    
149                    if vectors_per_cluster >= self.min_vectors_per_cluster {
150                        best_config = SubspaceConfig {
151                            effective_subspaces: num_subspaces,
152                            subspace_size,
153                            strategy: QuantizationStrategy::MultiSubspace,
154                            codebook_size_l1: self.original_codebook_l1,
155                            codebook_size_l2: self.original_codebook_l2,
156                            min_vectors_per_cluster: vectors_per_cluster,
157                        };
158                        break;
159                    }
160                }
161            }
162        }
163
164        best_config
165    }
166
167    /// Print configuration details for debugging
168    pub fn print_config(&self, config: &SubspaceConfig, rows: usize, cols: usize) {
169        match config.strategy {
170            QuantizationStrategy::ScalarQuantization => {
171                println!("🔧 NOVAQ: Using scalar quantization for {}×{} tensor (very small)", rows, cols);
172                println!("   Codebook size: {} entries", config.codebook_size_l1);
173            },
174            QuantizationStrategy::SingleSubspaceEnhanced => {
175                println!("🔧 NOVAQ: Using enhanced single-subspace for {}×{} tensor", rows, cols);
176                println!("   Enhanced codebook: L1={}, L2={}", config.codebook_size_l1, config.codebook_size_l2);
177            },
178            QuantizationStrategy::MultiSubspace => {
179                if config.effective_subspaces != self.original_subspaces {
180                    println!("🔧 NOVAQ: Adjusted subspaces from {} to {} for {}×{} tensor", 
181                             self.original_subspaces, config.effective_subspaces, rows, cols);
182                }
183                println!("   Subspace size: {}, Codebooks: L1={}, L2={}", 
184                         config.subspace_size, config.codebook_size_l1, config.codebook_size_l2);
185            },
186            QuantizationStrategy::UniformDithered => {
187                println!("🔧 NOVAQ: Using uniform dithered quantization for {}×{} tensor", rows, cols);
188            },
189        }
190    }
191}
192
193/// Fallback quantizer for edge cases where standard subspace quantization fails
194#[derive(Debug)]
195pub struct FallbackQuantizer {
196    stability_guard: NumericalStabilityGuard,
197}
198
199impl FallbackQuantizer {
200    pub fn new() -> Self {
201        Self {
202            stability_guard: NumericalStabilityGuard::default(),
203        }
204    }
205
206    /// Perform scalar quantization for very small tensors
207    pub fn scalar_quantize(
208        &mut self,
209        weights: &WeightMatrix,
210        config: &SubspaceConfig,
211    ) -> Result<(VectorCodebooks, QuantizationIndices)> {
212        let rows = weights.rows();
213        let cols = weights.cols();
214
215        // Create single codebook for entire vector
216        let mut vectors: Vec<Vec<f32>> = Vec::with_capacity(rows);
217        for row_idx in 0..rows {
218            let row = weights.get_row(row_idx).to_vec();
219            vectors.push(row);
220        }
221
222        // Simple k-means for single codebook
223        let l1_codebook = self.build_simple_codebook(&vectors, config.codebook_size_l1)?;
224        
225        // Minimal residual codebook (mostly zeros for scalar case)
226        let zero_vector = vec![0.0; cols];
227        let l2_codebook = vec![
228            CodebookEntry { centroid: zero_vector.clone(), usage_count: 0 },
229            CodebookEntry { centroid: zero_vector, usage_count: 0 },
230        ];
231
232        // Assign indices
233        let mut level1_indices = vec![vec![0u8]; rows];
234        let level2_indices = vec![vec![0u8]; rows]; // All zeros for residual
235
236        for (row_idx, vector) in vectors.iter().enumerate() {
237            let (nearest_idx, _) = self.find_nearest_centroid(vector, &l1_codebook);
238            level1_indices[row_idx][0] = nearest_idx as u8;
239        }
240
241        let codebooks = VectorCodebooks {
242            level1_codebooks: vec![l1_codebook],
243            level2_codebooks: vec![l2_codebook],
244            subspace_size: cols,
245        };
246
247        let indices = QuantizationIndices {
248            level1_indices,
249            level2_indices,
250        };
251
252        Ok((codebooks, indices))
253    }
254
255    /// Build a simple k-means codebook with stability protection
256    fn build_simple_codebook(
257        &mut self,
258        vectors: &[Vec<f32>],
259        k: usize,
260    ) -> Result<Vec<CodebookEntry>> {
261        if vectors.is_empty() {
262            return Err("Cannot build codebook from empty vector set".into());
263        }
264
265        let vector_dim = vectors[0].len();
266        let n_vectors = vectors.len();
267        let effective_k = k.min(n_vectors);
268
269        // Initialize centroids with actual data points to avoid extreme values
270        let mut centroids = Vec::with_capacity(effective_k);
271        for i in 0..effective_k {
272            let sample_idx = (i * n_vectors / effective_k).min(n_vectors - 1);
273            centroids.push(vectors[sample_idx].clone());
274        }
275
276        let mut assignments = vec![0; n_vectors];
277        let max_iterations = 50; // Reduced iterations for small tensors
278
279        // Simplified k-means with stability protection
280        for _iteration in 0..max_iterations {
281            // Assignment step
282            for (vec_idx, vector) in vectors.iter().enumerate() {
283                let (nearest_idx, _) = self.find_nearest_centroid_raw(vector, &centroids);
284                assignments[vec_idx] = nearest_idx;
285            }
286
287            // Update step with stability protection
288            for cluster_idx in 0..effective_k {
289                let cluster_vectors: Vec<&Vec<f32>> = vectors.iter()
290                    .enumerate()
291                    .filter(|(idx, _)| assignments[*idx] == cluster_idx)
292                    .map(|(_, vec)| vec)
293                    .collect();
294
295                if !cluster_vectors.is_empty() {
296                    for dim in 0..vector_dim {
297                        let values: Vec<f32> = cluster_vectors.iter().map(|v| v[dim]).collect();
298                        centroids[cluster_idx][dim] = self.stability_guard.safe_mean(&values);
299                    }
300                }
301            }
302        }
303
304        // Create codebook entries
305        let mut codebook = Vec::with_capacity(effective_k);
306        for cluster_idx in 0..effective_k {
307            let usage_count = assignments.iter().filter(|&&idx| idx == cluster_idx).count();
308            
309            // Ensure centroid is stable
310            self.stability_guard.sanitize_vector(&mut centroids[cluster_idx]);
311            
312            codebook.push(CodebookEntry {
313                centroid: centroids[cluster_idx].clone(),
314                usage_count,
315            });
316        }
317
318        Ok(codebook)
319    }
320
321    /// Find nearest centroid in raw centroid list
322    fn find_nearest_centroid_raw(&mut self, vector: &[f32], centroids: &[Vec<f32>]) -> (usize, f32) {
323        let mut min_distance = f32::INFINITY;
324        let mut nearest_idx = 0;
325
326        for (idx, centroid) in centroids.iter().enumerate() {
327            let distance = self.euclidean_distance(vector, centroid);
328            if distance < min_distance {
329                min_distance = distance;
330                nearest_idx = idx;
331            }
332        }
333
334        (nearest_idx, min_distance)
335    }
336
337    /// Find nearest centroid in codebook entries
338    fn find_nearest_centroid(&mut self, vector: &[f32], codebook: &[CodebookEntry]) -> (usize, f32) {
339        let centroids: Vec<Vec<f32>> = codebook.iter().map(|entry| entry.centroid.clone()).collect();
340        self.find_nearest_centroid_raw(vector, &centroids)
341    }
342
343    /// Calculate stable Euclidean distance
344    fn euclidean_distance(&mut self, a: &[f32], b: &[f32]) -> f32 {
345        let mut sum = 0.0;
346        for (&x, &y) in a.iter().zip(b.iter()) {
347            let diff = self.stability_guard.sanitize_value(x - y);
348            sum += diff * diff;
349        }
350        self.stability_guard.safe_sqrt(sum)
351    }
352}
353
354impl Default for FallbackQuantizer {
355    fn default() -> Self {
356        Self::new()
357    }
358}
359
360#[cfg(test)]
361mod tests {
362    use super::*;
363
364    #[test]
365    fn test_subspace_strategy_creation() {
366        let strategy = SubspaceStrategy::new(4, 16, 4);
367        assert_eq!(strategy.original_subspaces, 4);
368        assert_eq!(strategy.original_codebook_l1, 16);
369        assert_eq!(strategy.original_codebook_l2, 4);
370    }
371
372    #[test]
373    fn test_scalar_config_for_small_tensor() {
374        let strategy = SubspaceStrategy::new(4, 16, 4);
375        let config = strategy.determine_config(8, 1); // 8x1 tensor
376        
377        assert_eq!(config.strategy, QuantizationStrategy::ScalarQuantization);
378        assert_eq!(config.effective_subspaces, 1);
379        assert_eq!(config.subspace_size, 1);
380        assert!(config.codebook_size_l1 <= 8); // Should be reasonable for small tensor
381    }
382
383    #[test]
384    fn test_single_subspace_config() {
385        let strategy = SubspaceStrategy::new(4, 16, 4);
386        let config = strategy.determine_config(100, 3); // 100x3 tensor
387        
388        assert_eq!(config.strategy, QuantizationStrategy::SingleSubspaceEnhanced);
389        assert_eq!(config.effective_subspaces, 1);
390        assert_eq!(config.subspace_size, 3);
391        assert!(config.codebook_size_l1 >= 16); // Should be enhanced
392    }
393
394    #[test]
395    fn test_multi_subspace_config() {
396        let strategy = SubspaceStrategy::new(4, 16, 4);
397        let config = strategy.determine_config(1000, 16); // 1000x16 tensor
398        
399        assert_eq!(config.strategy, QuantizationStrategy::MultiSubspace);
400        assert!(config.effective_subspaces > 1);
401        assert_eq!(config.subspace_size, 16 / config.effective_subspaces);
402    }
403
404    #[test]
405    fn test_degenerate_cases() {
406        let strategy = SubspaceStrategy::new(4, 16, 4);
407        
408        // Zero rows
409        let config = strategy.determine_config(0, 10);
410        assert_eq!(config.strategy, QuantizationStrategy::ScalarQuantization);
411        
412        // Zero cols
413        let config = strategy.determine_config(10, 0);
414        assert_eq!(config.strategy, QuantizationStrategy::ScalarQuantization);
415    }
416
417    #[test]
418    fn test_fallback_quantizer() {
419        let mut quantizer = FallbackQuantizer::new();
420        let data = vec![1.0, 2.0, 3.0, 4.0];
421        let weights = WeightMatrix::new(data, vec![2, 2], "test".to_string());
422        
423        let config = SubspaceConfig {
424            effective_subspaces: 1,
425            subspace_size: 2,
426            strategy: QuantizationStrategy::ScalarQuantization,
427            codebook_size_l1: 2,
428            codebook_size_l2: 2,
429            min_vectors_per_cluster: 1,
430        };
431        
432        let result = quantizer.scalar_quantize(&weights, &config);
433        assert!(result.is_ok());
434        
435        let (codebooks, indices) = result.unwrap();
436        assert_eq!(codebooks.level1_codebooks.len(), 1);
437        assert_eq!(indices.level1_indices.len(), 2);
438    }
439}