Skip to main content

trustformers_wasm/optimization/quantization/
quantizer.rs

1//! Web-optimized quantizer with runtime adaptation
2
3use crate::optimization::quantization::algorithms::*;
4use crate::optimization::quantization::config::*;
5use serde::{Deserialize, Serialize};
6use std::vec::Vec;
7use wasm_bindgen::prelude::*;
8
9/// Result of quantization operation
10#[wasm_bindgen]
11#[derive(Debug, Clone)]
12pub struct QuantizationResult {
13    data: Vec<u8>,
14    stats: QuantizationStats,
15}
16
17#[wasm_bindgen]
18impl QuantizationResult {
19    /// Get the quantized data
20    pub fn data(&self) -> Vec<u8> {
21        self.data.clone()
22    }
23
24    /// Get a summary of the quantization
25    pub fn summary(&self) -> String {
26        format!(
27            "Quantization completed: {:.1}x compression, {:.1}% size reduction, {:.1}x estimated speedup",
28            self.stats.compression_ratio(),
29            self.stats.size_reduction_percent(),
30            self.stats.estimated_speedup()
31        )
32    }
33
34    /// Get detailed statistics
35    pub fn stats(&self) -> QuantizationStats {
36        self.stats.clone()
37    }
38}
39
40/// Web-optimized quantizer with runtime adaptation
41#[wasm_bindgen]
42pub struct WebQuantizer {
43    config: QuantizationConfig,
44    #[allow(dead_code)]
45    device_capabilities: DeviceCapabilities,
46    #[allow(dead_code)]
47    runtime_monitor: RuntimeMonitor,
48    adaptive_state: AdaptiveQuantizationState,
49}
50
51#[wasm_bindgen]
52impl WebQuantizer {
53    /// Create a new web quantizer
54    #[wasm_bindgen(constructor)]
55    pub fn new(config: QuantizationConfig) -> Self {
56        let device_capabilities = DeviceCapabilities {
57            supports_int8: true,
58            supports_int4: true,
59            supports_fp16: true,
60            memory_bandwidth_gb_s: 100.0,
61            compute_capability: ComputeCapability::Medium,
62        };
63
64        let runtime_monitor = RuntimeMonitor {
65            inference_times: Vec::new(),
66            memory_usage: Vec::new(),
67            accuracy_scores: Vec::new(),
68            thermal_state: ThermalState::Nominal,
69            adaptation_history: Vec::new(),
70        };
71
72        let adaptive_state = AdaptiveQuantizationState {
73            current_strategy: config.strategy(),
74            current_precision: config.precision(),
75            adaptation_rate: 0.1,
76            performance_target: config.performance_threshold(),
77            accuracy_target: config.accuracy_threshold(),
78            last_adaptation: 0.0,
79            confidence_score: 0.8,
80        };
81
82        Self {
83            config,
84            device_capabilities,
85            runtime_monitor,
86            adaptive_state,
87        }
88    }
89
90    /// Quantize tensor data using the configured strategy
91    pub fn quantize(&self, data: &[f32]) -> Result<Vec<f32>, JsValue> {
92        match self.adaptive_state.current_strategy {
93            QuantizationStrategy::None => Ok(data.to_vec()),
94            QuantizationStrategy::Dynamic => {
95                apply_dynamic_quantization(data, self.adaptive_state.current_precision)
96            },
97            QuantizationStrategy::Static => {
98                apply_static_quantization(data, self.adaptive_state.current_precision)
99            },
100            QuantizationStrategy::PostTraining => {
101                apply_post_training_quantization(data, self.adaptive_state.current_precision)
102            },
103            QuantizationStrategy::AWQ => {
104                apply_awq_quantization(data, self.adaptive_state.current_precision)
105            },
106            QuantizationStrategy::GPTQ => {
107                apply_gptq_quantization(data, self.adaptive_state.current_precision)
108            },
109            QuantizationStrategy::SmoothQuant => {
110                apply_smoothquant_quantization(data, self.adaptive_state.current_precision)
111            },
112            QuantizationStrategy::LLMInt8 => {
113                apply_llm_int8_quantization(data, self.adaptive_state.current_precision)
114            },
115            QuantizationStrategy::QLoRA => {
116                apply_qlora_quantization(data, self.adaptive_state.current_precision)
117            },
118            QuantizationStrategy::GGML => {
119                apply_ggml_quantization(data, self.adaptive_state.current_precision)
120            },
121            QuantizationStrategy::AdaptiveBitwidth => {
122                apply_adaptive_bitwidth_quantization(data, self.adaptive_state.current_precision)
123            },
124            QuantizationStrategy::OutlierAware => {
125                apply_outlier_aware_quantization(data, self.adaptive_state.current_precision)
126            },
127            QuantizationStrategy::HQQ => {
128                apply_hqq_quantization(data, self.adaptive_state.current_precision)
129            },
130            QuantizationStrategy::SpQR => {
131                apply_spqr_quantization(data, self.adaptive_state.current_precision)
132            },
133            QuantizationStrategy::AQLM => {
134                apply_aqlm_quantization(data, self.adaptive_state.current_precision)
135            },
136            _ => Err(JsValue::from_str("Unsupported quantization strategy")),
137        }
138    }
139
140    /// Get quantization statistics
141    pub fn get_stats(&self, original_data: &[f32], quantized_data: &[f32]) -> QuantizationStats {
142        let bytes_per_element = self.adaptive_state.current_precision.bytes_per_element();
143        let original_size = original_data.len() * 4; // 4 bytes per f32, always — original is raw f32
144        let quantized_size = (quantized_data.len() as f32 * bytes_per_element).ceil() as usize;
145        let compression_ratio = original_size as f32 / quantized_size as f32;
146        let size_reduction = (1.0 - quantized_size as f32 / original_size as f32) * 100.0;
147
148        QuantizationStats::new(
149            original_size,
150            quantized_size,
151            compression_ratio,
152            size_reduction,
153            compression_ratio * 0.8, // Simplified estimation
154            self.adaptive_state.current_strategy,
155            self.adaptive_state.current_precision,
156        )
157    }
158
159    /// Check if a model should be quantized based on size
160    pub fn should_quantize(&self, model_size_bytes: usize) -> bool {
161        let model_size_mb = model_size_bytes as f32 / (1024.0 * 1024.0);
162        model_size_mb > self.config.target_size_mb()
163    }
164
165    /// Quantize model data
166    pub fn quantize_model(&self, model_data: &[u8]) -> Result<QuantizationResult, JsValue> {
167        // Convert bytes to f32 for processing
168        let float_data: Vec<f32> = model_data
169            .chunks_exact(4)
170            .map(|chunk| f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]))
171            .collect();
172
173        // Apply quantization
174        let quantized_floats = self.quantize(&float_data)?;
175
176        // Get stats before converting to bytes
177        let stats = self.get_stats(&float_data, &quantized_floats);
178
179        // Convert back to bytes
180        let quantized_bytes: Vec<u8> =
181            quantized_floats.into_iter().flat_map(|f| f.to_le_bytes()).collect();
182
183        Ok(QuantizationResult {
184            data: quantized_bytes,
185            stats,
186        })
187    }
188
189    /// Get recommended quantization settings for a given model size
190    pub fn get_recommended_settings(&self, model_size_bytes: usize) -> QuantizationConfig {
191        let model_size_mb = model_size_bytes as f32 / (1024.0 * 1024.0);
192
193        if model_size_mb < 10.0 {
194            QuantizationConfig::new(QuantizationStrategy::None, QuantizationPrecision::FP16)
195        } else if model_size_mb < 50.0 {
196            QuantizationConfig::new(QuantizationStrategy::Dynamic, QuantizationPrecision::FP16)
197        } else if model_size_mb < 200.0 {
198            QuantizationConfig::new(
199                QuantizationStrategy::PostTraining,
200                QuantizationPrecision::INT8,
201            )
202        } else {
203            QuantizationConfig::new(QuantizationStrategy::AWQ, QuantizationPrecision::INT4)
204        }
205    }
206}
207
208/// Quantized model data
209#[derive(Debug, Clone, Serialize, Deserialize)]
210pub struct QuantizedModelData {
211    pub quantized_weights: Vec<Vec<f32>>,
212    pub scale_factors: Vec<f32>,
213    pub zero_points: Vec<f32>,
214    pub metadata: QuantizationMetadata,
215}
216
217/// Quantization metadata
218#[derive(Debug, Clone, Serialize, Deserialize)]
219pub struct QuantizationMetadata {
220    pub strategy: QuantizationStrategy,
221    pub precision: QuantizationPrecision,
222    pub compression_ratio: f32,
223    pub accuracy_retention: f32,
224}
225
226impl QuantizedModelData {
227    pub fn new(strategy: QuantizationStrategy, precision: QuantizationPrecision) -> Self {
228        Self {
229            quantized_weights: Vec::new(),
230            scale_factors: Vec::new(),
231            zero_points: Vec::new(),
232            metadata: QuantizationMetadata {
233                strategy,
234                precision,
235                compression_ratio: 1.0,
236                accuracy_retention: 1.0,
237            },
238        }
239    }
240}
241
242#[cfg(test)]
243mod tests {
244    use super::*;
245
246    #[test]
247    fn test_quantizer_creation() {
248        let config = QuantizationConfig::auto();
249        let quantizer = WebQuantizer::new(config);
250        assert_eq!(
251            quantizer.adaptive_state.current_strategy,
252            QuantizationStrategy::Dynamic
253        );
254    }
255
256    #[test]
257    fn test_basic_quantization() {
258        let config =
259            QuantizationConfig::new(QuantizationStrategy::Dynamic, QuantizationPrecision::INT8);
260        let quantizer = WebQuantizer::new(config);
261        let data = vec![1.0, 2.0, 3.0, 4.0];
262        let result = quantizer.quantize(&data);
263        assert!(result.is_ok());
264    }
265
266    #[test]
267    fn test_quantization_stats() {
268        let config = QuantizationConfig::auto();
269        let quantizer = WebQuantizer::new(config);
270        let original = vec![1.0, 2.0, 3.0, 4.0];
271        let quantized = vec![0.5, 1.0, 1.5, 2.0];
272        let stats = quantizer.get_stats(&original, &quantized);
273        assert!(stats.compression_ratio() >= 1.0);
274    }
275
276    #[test]
277    fn test_quantized_model_data() {
278        let data = QuantizedModelData::new(QuantizationStrategy::AWQ, QuantizationPrecision::INT8);
279        assert_eq!(data.metadata.strategy, QuantizationStrategy::AWQ);
280        assert_eq!(data.metadata.precision, QuantizationPrecision::INT8);
281    }
282
283    #[test]
284    fn test_get_stats_is_bit_width_aware_for_int4() {
285        let config =
286            QuantizationConfig::new(QuantizationStrategy::Dynamic, QuantizationPrecision::INT4);
287        let quantizer = WebQuantizer::new(config);
288        let original = vec![1.0f32; 100];
289        let quantized = vec![0.5f32; 100]; // same length as original, matching real apply_* output
290        let stats = quantizer.get_stats(&original, &quantized);
291
292        assert_eq!(stats.original_size_bytes(), 400); // 100 * 4 bytes, unaffected by target precision
293        assert_eq!(stats.quantized_size_bytes(), 50); // 100 * 0.5 bytes (INT4) = 50, not 400 like before the fix
294        assert!(stats.compression_ratio() > 4.0); // real ~8x compression, far more than the old fixed 1.0x
295    }
296
297    #[test]
298    fn test_get_stats_scales_per_precision_not_just_a_new_fixed_constant() {
299        let config =
300            QuantizationConfig::new(QuantizationStrategy::Dynamic, QuantizationPrecision::INT8);
301        let quantizer = WebQuantizer::new(config);
302        let original = vec![1.0f32; 100];
303        let quantized = vec![0.5f32; 100];
304        let stats = quantizer.get_stats(&original, &quantized);
305        assert_eq!(stats.quantized_size_bytes(), 100); // INT8 = 1 byte/element, different from INT4's 50
306    }
307}