1use crate::Result;
2use super::{WeightMatrix, VectorCodebooks, QuantizationIndices, CodebookEntry, SubspaceStrategy, FallbackQuantizer, QuantizationStrategy};
3use rand::{Rng, SeedableRng};
4use rand_chacha::ChaCha8Rng;
5
6#[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 pub fn build_codebooks(&mut self, weights: &WeightMatrix) -> Result<(VectorCodebooks, QuantizationIndices)> {
43 let rows = weights.rows();
44 let cols = weights.cols();
45
46 let config = self.subspace_strategy.determine_config(rows, cols);
48
49 self.subspace_strategy.print_config(&config, rows, cols);
51
52 if config.strategy == QuantizationStrategy::ScalarQuantization {
54 return self.fallback_quantizer.scalar_quantize(weights, &config);
55 }
56
57 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 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 if end_col > cols {
75 return Err(format!("Subspace bounds error: end_col {} > cols {}", end_col, cols).into());
76 }
77
78 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 let l1_codebook = self.build_kmeans_codebook(&subspace_vectors, l1_size)?;
91
92 let mut residuals = Vec::with_capacity(rows);
94 for (row_idx, vector) in subspace_vectors.iter().enumerate() {
95 let (nearest_idx, _) = self.find_nearest_centroid(vector, &l1_codebook);
97 level1_indices[row_idx][subspace_idx] = nearest_idx as u8;
98
99 let residual = self.calculate_residual(vector, &l1_codebook[nearest_idx].centroid);
101 residuals.push(residual);
102 }
103
104 let l2_codebook = self.build_kmeans_codebook(&residuals, l2_size)?;
106
107 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 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 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 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 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 let effective_k = k.min(n_vectors);
176
177 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 for iteration in 0..self.max_kmeans_iterations {
192 for (vec_idx, vector) in vectors.iter().enumerate() {
194 let (nearest_idx, _) = self.find_nearest_centroid_raw(vector, ¢roids);
195 assignments[vec_idx] = nearest_idx;
196 }
197
198 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 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 for dim in 0..vector_dim {
215 centroids[cluster_idx][dim] = self.rng.gen_range(-1.0..1.0);
216 }
217 }
218 }
219
220 let converged = self.check_convergence(&prev_centroids, ¢roids);
222 if converged {
223 break;
224 }
225
226 prev_centroids = centroids.clone();
227 }
228
229 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 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, ¢roids)
246 }
247
248 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 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 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 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
286pub 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 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, 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 let metrics = QuantizationMetrics::calculate(&data, &reconstructed);
392 println!("Reconstruction MSE: {:.6}", metrics.mse);
393
394 assert!(metrics.mse < 10.0); }
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, ¢roid);
417 let expected = vec![3.0, 4.0, 5.0];
418
419 assert_eq!(residual, expected);
420 }
421}