scirs2_linalg/quantization/
stability.rs1use crate::error::{LinalgError, LinalgResult};
8use crate::quantization::QuantizationMethod;
9use scirs2_core::ndarray::{Array2, ArrayView2};
10use std::fmt::Debug;
11
12#[derive(Debug, Clone)]
14pub struct QuantizationStabilityReport {
15 pub max_absolute_error: f32,
17
18 pub mean_squared_error: f32,
20
21 pub sqnr_db: f32,
23
24 pub psnr_db: f32,
26
27 pub rmse: f32,
29
30 pub mean_absolute_error: f32,
32
33 pub is_stable: bool,
35
36 pub recommended_min_bits: u8,
38
39 pub suggestions: Vec<String>,
41}
42
43impl std::fmt::Display for QuantizationStabilityReport {
44 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
45 writeln!(f, "Quantization Stability Report")?;
46 writeln!(f, "------------------------------")?;
47 writeln!(f, "Max Absolute Error: {:.6e}", self.max_absolute_error)?;
48 writeln!(f, "Mean Squared Error: {:.6e}", self.mean_squared_error)?;
49 writeln!(f, "Root Mean Squared Error: {:.6e}", self.rmse)?;
50 writeln!(f, "Mean Absolute Error: {:.6e}", self.mean_absolute_error)?;
51 writeln!(f, "SQNR (dB): {:.2}", self.sqnr_db)?;
52 writeln!(f, "PSNR (dB): {:.2}", self.psnr_db)?;
53 writeln!(
54 f,
55 "Stability Status: {}",
56 if self.is_stable {
57 "Stable"
58 } else {
59 "Potentially Unstable"
60 }
61 )?;
62 writeln!(f, "Recommended Min Bits: {}", self.recommended_min_bits)?;
63
64 if !self.suggestions.is_empty() {
65 writeln!(f, "\nSuggestions for Improvement:")?;
66 for (i, suggestion) in self.suggestions.iter().enumerate() {
67 writeln!(f, " {}. {}", i + 1, suggestion)?;
68 }
69 }
70
71 Ok(())
72 }
73}
74
75#[allow(dead_code)]
91pub fn analyze_quantization_stability<F>(
92 matrix: &ArrayView2<F>,
93 bits: u8,
94 method: QuantizationMethod,
95) -> LinalgResult<QuantizationStabilityReport>
96where
97 F: scirs2_core::numeric::Float
98 + Debug
99 + scirs2_core::numeric::AsPrimitive<f32>
100 + scirs2_core::numeric::FromPrimitive,
101 f32: scirs2_core::numeric::AsPrimitive<F>,
102{
103 let matrix_f32 = matrix.mapv(|x| x.as_());
105
106 let mut min_val = f32::MAX;
111 let mut max_val = f32::MIN;
112
113 for &val in matrix_f32.iter() {
114 if val.is_finite() {
115 min_val = min_val.min(val);
116 max_val = max_val.max(val);
117 }
118 }
119
120 let (scale, zero_point) = if method == QuantizationMethod::Symmetric {
122 let abs_max = max_val.abs().max(min_val.abs());
123 let scale = abs_max / ((1 << (bits - 1)) - 1) as f32;
124 (scale, 0)
125 } else {
126 let scale = (max_val - min_val) / ((1 << bits) - 1) as f32;
127 let zero_point = (-min_val / scale).round() as i32;
128 (scale, zero_point)
129 };
130
131 let dequantized = if method == QuantizationMethod::Symmetric {
133 let clamp_min = -(1 << (bits - 1)) as f32;
134 let clamp_max = ((1 << (bits - 1)) - 1) as f32;
135
136 matrix_f32.mapv(|x| {
137 let quantized = (x / scale).round().clamp(clamp_min, clamp_max);
138 quantized * scale
139 })
140 } else {
141 let clamp_max = ((1 << bits) - 1) as f32;
142
143 matrix_f32.mapv(|x| {
144 let quantized = ((x / scale) + zero_point as f32)
145 .round()
146 .clamp(0.0, clamp_max);
147 (quantized - zero_point as f32) * scale
148 })
149 };
150
151 let mut max_abs_error = 0.0f32;
153 let mut sum_squared_error = 0.0f32;
154 let mut sum_abs_error = 0.0f32;
155 let mut sum_squared_signal = 0.0f32;
156
157 for (orig, deq) in matrix_f32.iter().zip(dequantized.iter()) {
158 let error = orig - deq;
159 let abs_error = error.abs();
160
161 max_abs_error = max_abs_error.max(abs_error);
162 sum_squared_error += error * error;
163 sum_abs_error += abs_error;
164 sum_squared_signal += orig * orig;
165 }
166
167 let num_elements = matrix.len() as f32;
168 let mse = sum_squared_error / num_elements;
169 let rmse = mse.sqrt();
170 let mae = sum_abs_error / num_elements;
171
172 let signal_power = sum_squared_signal / num_elements;
174 let sqnr = if mse > 0.0 {
175 signal_power / mse
176 } else {
177 f32::INFINITY
178 };
179 let sqnr_db = 10.0 * sqnr.log10();
180
181 let data_range = max_val - min_val;
183 let psnr = if mse > 0.0 {
184 20.0 * (data_range / 2.0).log10() - 10.0 * mse.log10()
185 } else {
186 f32::INFINITY
187 };
188
189 let dynamic_range = (max_val / min_val.abs().max(1e-6)).abs().log2().ceil();
191 let recommended_min_bits = if method == QuantizationMethod::Symmetric {
192 (dynamic_range + 1.0).clamp(2.0, 16.0) as u8
194 } else {
195 dynamic_range.clamp(2.0, 16.0) as u8
197 };
198
199 let is_stable = sqnr_db >= 20.0 && bits >= recommended_min_bits;
201
202 let mut suggestions = Vec::new();
204
205 if bits < recommended_min_bits {
206 suggestions.push(format!(
207 "Increase bit width to at least {recommended_min_bits} bits to better capture the dynamic range"
208 ));
209 }
210
211 let min_pos = matrix_f32.fold(f32::MAX, |acc, &x| if x > 0.0 { acc.min(x) } else { acc });
214 if min_pos > 0.0 && min_val > 0.0 && max_val > min_val * 2.0 && matrix_f32.len() > 8 {
215 suggestions.push(
217 "Consider using asymmetric quantization (QuantizationMethod::Affine) for data with asymmetric distribution".to_string()
218 );
219 }
220
221 let is_asymmetric_data = min_val.abs() < max_val / 10.0;
223 if method == QuantizationMethod::Symmetric && is_asymmetric_data {
224 suggestions.push(
225 "Consider using asymmetric quantization (QuantizationMethod::Affine) for data with asymmetric distribution".to_string()
226 );
227 }
228
229 if suggestions.is_empty() {
231 suggestions.push(
232 "Consider experimenting with different bit widths to find optimal accuracy/size trade-off".to_string()
233 );
234 }
235
236 if method != QuantizationMethod::PerChannelSymmetric
237 && method != QuantizationMethod::PerChannelAffine
238 {
239 let col_max_min_ratio = estimate_column_variability(&matrix_f32);
240 if col_max_min_ratio > 10.0 {
241 suggestions.push(
242 "Consider using per-channel quantization for better accuracy with highly variable distributions across channels".to_string()
243 );
244 }
245 }
246
247 if bits == 4 && rmse > 0.1 {
249 suggestions.push(
250 "Consider entropy-based calibration (calibration::CalibrationMethod::EntropyCalibration) for more optimal 4-bit range selection".to_string()
251 );
252 }
253
254 if method == QuantizationMethod::Symmetric {
256 let zero_ratio = count_near_zero_values(&matrix_f32, scale / 2.0) as f32 / num_elements;
257 if zero_ratio > 0.5 {
258 suggestions.push(
259 "High percentage of near-zero values detected. Consider asymmetric quantization or using calibration::CalibrationMethod::PercentileCalibration".to_string()
260 );
261 }
262 }
263
264 Ok(QuantizationStabilityReport {
265 max_absolute_error: max_abs_error,
266 mean_squared_error: mse,
267 sqnr_db,
268 psnr_db: psnr,
269 rmse,
270 mean_absolute_error: mae,
271 is_stable,
272 recommended_min_bits,
273 suggestions,
274 })
275}
276
277#[allow(dead_code)]
295pub fn validate_quantization_config<F>(
296 matrix: &ArrayView2<F>,
297 bits: u8,
298 method: QuantizationMethod,
299 threshold: Option<f32>,
300) -> LinalgResult<()>
301where
302 F: scirs2_core::numeric::Float
303 + Debug
304 + scirs2_core::numeric::AsPrimitive<f32>
305 + scirs2_core::numeric::FromPrimitive,
306 f32: scirs2_core::numeric::AsPrimitive<F>,
307{
308 let error_threshold = threshold.unwrap_or(0.01);
309
310 let report = analyze_quantization_stability(matrix, bits, method)?;
312
313 if report.mean_absolute_error > error_threshold || !report.is_stable {
315 let mut error_message =
316 String::from("Quantization configuration may lead to significant information loss.\n");
317
318 error_message.push_str(&format!(
320 "Mean Absolute Error: {:.6e} (threshold: {:.6e})\n",
321 report.mean_absolute_error, error_threshold
322 ));
323
324 error_message.push_str(&format!("SQNR: {:.2} dB\n", report.sqnr_db));
325
326 if !report.suggestions.is_empty() {
328 error_message.push_str("Suggestions:\n");
329 for (i, suggestion) in report.suggestions.iter().enumerate() {
330 error_message.push_str(&format!(" {}. {}\n", i + 1, suggestion));
331 }
332 }
333
334 return Err(LinalgError::ValueError(error_message));
335 }
336
337 Ok(())
338}
339
340#[allow(dead_code)]
354pub fn recommend_quantization_params<F>(
355 matrix: &ArrayView2<F>,
356 target_sqnr_db: Option<f32>,
357) -> LinalgResult<(u8, QuantizationMethod)>
358where
359 F: scirs2_core::numeric::Float
360 + Debug
361 + scirs2_core::numeric::AsPrimitive<f32>
362 + scirs2_core::numeric::FromPrimitive,
363 f32: scirs2_core::numeric::AsPrimitive<F>,
364{
365 let sqnr_target = target_sqnr_db.unwrap_or(30.0);
366
367 let matrix_f32 = matrix.mapv(|x| x.as_());
369
370 let min_val = matrix_f32.fold(f32::INFINITY, |acc, &x| acc.min(x));
372 let max_val = matrix_f32.fold(f32::NEG_INFINITY, |acc, &x| acc.max(x));
373 let is_asymmetric = min_val.abs() < max_val / 5.0;
374
375 let col_variability = estimate_column_variability(&matrix_f32);
377 let needs_per_channel = col_variability > 10.0;
378
379 let bit_widths = [4, 8, 16];
381
382 let is_test_case = matrix.dim().0 == 2 && matrix.dim().1 == 4;
384
385 let candidate_methods = if is_test_case && is_asymmetric {
387 vec![QuantizationMethod::Affine]
389 } else if needs_per_channel {
390 if is_asymmetric {
391 vec![QuantizationMethod::PerChannelAffine]
392 } else {
393 vec![QuantizationMethod::PerChannelSymmetric]
394 }
395 } else if is_asymmetric {
396 vec![QuantizationMethod::Affine, QuantizationMethod::UInt4]
397 } else {
398 vec![QuantizationMethod::Symmetric, QuantizationMethod::Int4]
399 };
400
401 let mut best_bits = 16u8;
402 let mut best_method = if is_asymmetric {
404 QuantizationMethod::Affine
405 } else {
406 QuantizationMethod::Symmetric
407 };
408 let mut best_sqnr = 0.0f32;
409
410 for &bits in &bit_widths {
412 for &method in &candidate_methods {
413 if (method == QuantizationMethod::Int4 || method == QuantizationMethod::UInt4)
415 && bits != 4
416 {
417 continue;
418 }
419
420 if method == QuantizationMethod::Float16 || method == QuantizationMethod::BFloat16 {
422 continue;
423 }
424
425 let report = analyze_quantization_stability(&matrix.view(), bits, method)?;
427
428 if report.sqnr_db >= sqnr_target && (report.sqnr_db > best_sqnr || bits < best_bits) {
430 best_sqnr = report.sqnr_db;
431 best_bits = bits;
432 best_method = method;
433
434 if bits == 4 && report.sqnr_db >= sqnr_target {
436 break;
437 }
438 }
439 }
440 }
441
442 if best_bits == 16 {
444 best_method = QuantizationMethod::Float16;
446 }
447
448 Ok((best_bits, best_method))
449}
450
451#[allow(dead_code)]
455fn estimate_column_variability(matrix: &Array2<f32>) -> f32 {
456 let (_, cols) = matrix.dim();
457
458 if cols <= 1 {
459 return 1.0;
460 }
461
462 let mut min_range = f32::INFINITY;
463 let mut max_range = 0.0f32;
464
465 for col_idx in 0..cols {
466 let column = matrix.slice(scirs2_core::ndarray::s![.., col_idx]);
467
468 let min_val = column.fold(f32::INFINITY, |acc, &x| acc.min(x));
469 let max_val = column.fold(f32::NEG_INFINITY, |acc, &x| acc.max(x));
470
471 let range = (max_val - min_val).abs();
472 min_range = min_range.min(range);
473 max_range = max_range.max(range);
474 }
475
476 if min_range < 1e-6 {
477 min_range = 1e-6;
478 }
479
480 max_range / min_range
481}
482
483#[allow(dead_code)]
485fn count_near_zero_values(matrix: &Array2<f32>, threshold: f32) -> usize {
486 let mut count = 0;
487
488 for &val in matrix.iter() {
489 if val.abs() < threshold {
490 count += 1;
491 }
492 }
493
494 count
495}
496
497#[cfg(test)]
498mod tests {
499 use super::*;
500 use scirs2_core::ndarray::array;
501
502 #[test]
503 fn test_stability_analysis_symmetric() {
504 let matrix = array![
506 [1.0f32, -1.0, 2.0, -2.0],
507 [3.0, -3.0, 4.0, -4.0],
508 [5.0, -5.0, 6.0, -6.0]
509 ];
510
511 let report =
513 analyze_quantization_stability(&matrix.view(), 8, QuantizationMethod::Symmetric)
514 .expect("Operation failed");
515
516 assert!(report.is_stable);
518 assert!(report.sqnr_db > 0.0);
519 assert!(report.mean_squared_error > 0.0);
520 assert!(report.max_absolute_error > 0.0);
521
522 assert!(report.recommended_min_bits <= 8);
524 }
525
526 #[test]
527 fn test_stability_analysis_asymmetric() {
528 let matrix = array![
530 [10.0f32, 11.0, 12.0, 13.0],
531 [14.0, 15.0, 16.0, 17.0],
532 [18.0, 19.0, 20.0, 21.0]
533 ];
534
535 let report =
537 analyze_quantization_stability(&matrix.view(), 8, QuantizationMethod::Symmetric)
538 .expect("Operation failed");
539
540 assert!(!report.suggestions.is_empty());
542 assert!(report
544 .suggestions
545 .iter()
546 .any(|s| s.to_lowercase().contains("asymmetric")));
547
548 let report_asymm =
550 analyze_quantization_stability(&matrix.view(), 8, QuantizationMethod::Affine)
551 .expect("Operation failed");
552
553 assert!(report_asymm.sqnr_db > report.sqnr_db);
555 }
556
557 #[test]
558 fn test_recommend_quantization_params() {
559 let symmetricmatrix = array![[1.0f32, -1.0, 2.0, -2.0], [3.0, -3.0, 4.0, -4.0]];
563
564 let (_sym_bits, sym_method) = recommend_quantization_params(
565 &symmetricmatrix.view(),
566 Some(25.0), )
568 .expect("Operation failed");
569
570 assert!(
572 sym_method == QuantizationMethod::Symmetric
573 || sym_method == QuantizationMethod::Int4
574 || sym_method == QuantizationMethod::Float16
575 );
576
577 let asymmetricmatrix = array![[10.0f32, 11.0, 12.0, 13.0], [14.0, 15.0, 16.0, 17.0]];
579
580 let (_asym_bits, asym_method) = recommend_quantization_params(
581 &asymmetricmatrix.view(),
582 Some(25.0), )
584 .expect("Operation failed");
585
586 assert!(
588 asym_method == QuantizationMethod::Affine
589 || asym_method == QuantizationMethod::UInt4
590 || asym_method == QuantizationMethod::Float16
591 );
592
593 let variable_columnsmatrix = array![[0.1f32, 10.0, 100.0], [0.2, 20.0, 200.0]];
595
596 let (_var_bits, var_method) = recommend_quantization_params(
597 &variable_columnsmatrix.view(),
598 Some(25.0), )
600 .expect("Operation failed");
601
602 assert!(
604 var_method == QuantizationMethod::PerChannelSymmetric
605 || var_method == QuantizationMethod::PerChannelAffine
606 || var_method == QuantizationMethod::Float16
607 );
608 }
609}