Skip to main content

torsh_optim/
low_precision.rs

1use crate::OptimizerError;
2use std::collections::HashMap;
3use std::fmt;
4
5/// Trait for types that can be converted to/from low-precision representation
6pub trait LowPrecisionConvertible: Clone + fmt::Debug {
7    /// Convert to low-precision representation
8    fn to_low_precision(&self) -> LowPrecisionState;
9
10    /// Convert from low-precision representation
11    fn from_low_precision(state: &LowPrecisionState) -> Result<Self, OptimizerError>;
12}
13
14/// Low-precision state representation for memory-efficient optimizer states
15#[derive(Clone, Debug)]
16pub enum LowPrecisionState {
17    /// 16-bit float representation
18    F16(Vec<half::f16>),
19    /// 16-bit brain float representation
20    BF16(Vec<half::bf16>),
21    /// 8-bit integer representation with scale factor
22    I8 { values: Vec<i8>, scale: f32 },
23    /// 4-bit integer representation with scale factor
24    I4 { values: Vec<u8>, scale: f32 }, // packed 2 values per byte
25    /// Sparse representation for mostly-zero states
26    Sparse {
27        indices: Vec<usize>,
28        values: Vec<f32>,
29    },
30}
31
32impl LowPrecisionState {
33    /// Get the memory footprint in bytes
34    pub fn memory_footprint(&self) -> usize {
35        match self {
36            LowPrecisionState::F16(values) => values.len() * 2,
37            LowPrecisionState::BF16(values) => values.len() * 2,
38            LowPrecisionState::I8 { values, .. } => values.len() + 4, // +4 for scale
39            LowPrecisionState::I4 { values, .. } => (values.len() + 1) / 2 + 4, // packed + scale
40            LowPrecisionState::Sparse { indices, values } => {
41                indices.len() * 8 + values.len() * 4 // usize + f32
42            }
43        }
44    }
45
46    /// Convert to full precision f32 vector
47    pub fn to_f32(&self) -> Vec<f32> {
48        match self {
49            LowPrecisionState::F16(values) => values.iter().map(|&x| x.to_f32()).collect(),
50            LowPrecisionState::BF16(values) => values.iter().map(|&x| x.to_f32()).collect(),
51            LowPrecisionState::I8 { values, scale } => {
52                values.iter().map(|&x| (x as f32) * scale).collect()
53            }
54            LowPrecisionState::I4 { values, scale } => {
55                let mut result = Vec::with_capacity(values.len() * 2);
56                for &packed in values {
57                    let low = ((packed & 0x0F) as i8 - 8) as f32 * scale;
58                    let high = (((packed >> 4) as i8) - 8) as f32 * scale;
59                    result.push(low);
60                    result.push(high);
61                }
62                result
63            }
64            LowPrecisionState::Sparse { indices, values } => {
65                let max_idx = indices.iter().max().copied().unwrap_or(0);
66                let mut result = vec![0.0f32; max_idx + 1];
67                for (&idx, &val) in indices.iter().zip(values.iter()) {
68                    result[idx] = val;
69                }
70                result
71            }
72        }
73    }
74
75    /// Create from f32 vector with specified precision
76    pub fn from_f32(values: &[f32], precision: PrecisionType) -> Self {
77        match precision {
78            PrecisionType::F16 => {
79                let converted: Vec<half::f16> =
80                    values.iter().map(|&x| half::f16::from_f32(x)).collect();
81                LowPrecisionState::F16(converted)
82            }
83            PrecisionType::BF16 => {
84                let converted: Vec<half::bf16> =
85                    values.iter().map(|&x| half::bf16::from_f32(x)).collect();
86                LowPrecisionState::BF16(converted)
87            }
88            PrecisionType::I8 => {
89                let max_val = values.iter().fold(0.0f32, |acc, &x| acc.max(x.abs()));
90                let scale = max_val / 127.0;
91                let converted: Vec<i8> =
92                    values.iter().map(|&x| (x / scale).round() as i8).collect();
93                LowPrecisionState::I8 {
94                    values: converted,
95                    scale,
96                }
97            }
98            PrecisionType::I4 => {
99                let max_val = values.iter().fold(0.0f32, |acc, &x| acc.max(x.abs()));
100                let scale = max_val / 7.0;
101                let mut converted = Vec::with_capacity((values.len() + 1) / 2);
102
103                for chunk in values.chunks(2) {
104                    let low = ((chunk[0] / scale).round() as i8 + 8) as u8 & 0x0F;
105                    let high = if chunk.len() > 1 {
106                        (((chunk[1] / scale).round() as i8 + 8) as u8 & 0x0F) << 4
107                    } else {
108                        0
109                    };
110                    converted.push(low | high);
111                }
112
113                LowPrecisionState::I4 {
114                    values: converted,
115                    scale,
116                }
117            }
118            PrecisionType::Sparse(threshold) => {
119                let mut indices = Vec::new();
120                let mut sparse_values = Vec::new();
121
122                for (i, &val) in values.iter().enumerate() {
123                    if val.abs() > threshold {
124                        indices.push(i);
125                        sparse_values.push(val);
126                    }
127                }
128
129                LowPrecisionState::Sparse {
130                    indices,
131                    values: sparse_values,
132                }
133            }
134        }
135    }
136}
137
138/// Precision type for low-precision states
139#[derive(Clone, Debug)]
140pub enum PrecisionType {
141    /// 16-bit float
142    F16,
143    /// 16-bit brain float
144    BF16,
145    /// 8-bit integer with scale
146    I8,
147    /// 4-bit integer with scale
148    I4,
149    /// Sparse representation with threshold
150    Sparse(f32),
151}
152
153/// Low-precision optimizer wrapper
154pub struct LowPrecisionOptimizer<T> {
155    inner: T,
156    precision: PrecisionType,
157    state_cache: HashMap<String, LowPrecisionState>,
158}
159
160impl<T> LowPrecisionOptimizer<T> {
161    /// Create a new low-precision optimizer wrapper
162    pub fn new(inner: T, precision: PrecisionType) -> Self {
163        Self {
164            inner,
165            precision,
166            state_cache: HashMap::new(),
167        }
168    }
169
170    /// Get the precision type
171    pub fn precision(&self) -> &PrecisionType {
172        &self.precision
173    }
174
175    /// Get the inner optimizer
176    pub fn inner(&self) -> &T {
177        &self.inner
178    }
179
180    /// Get the inner optimizer mutably
181    pub fn inner_mut(&mut self) -> &mut T {
182        &mut self.inner
183    }
184
185    /// Store state in low precision
186    pub fn store_state(&mut self, key: String, values: &[f32]) {
187        let low_precision_state = LowPrecisionState::from_f32(values, self.precision.clone());
188        self.state_cache.insert(key, low_precision_state);
189    }
190
191    /// Load state from low precision
192    pub fn load_state(&self, key: &str) -> Option<Vec<f32>> {
193        self.state_cache.get(key).map(|state| state.to_f32())
194    }
195
196    /// Get total memory footprint of stored states
197    pub fn memory_footprint(&self) -> usize {
198        self.state_cache
199            .values()
200            .map(|state| state.memory_footprint())
201            .sum()
202    }
203
204    /// Get compression ratio compared to full precision
205    pub fn compression_ratio(&self) -> f32 {
206        let total_elements: usize = self
207            .state_cache
208            .values()
209            .map(|state| state.to_f32().len())
210            .sum();
211
212        if total_elements == 0 {
213            return 1.0;
214        }
215
216        let full_precision_size = total_elements * 4; // 4 bytes per f32
217        let compressed_size = self.memory_footprint();
218
219        full_precision_size as f32 / compressed_size as f32
220    }
221
222    /// Clear all cached states
223    pub fn clear_cache(&mut self) {
224        self.state_cache.clear();
225    }
226
227    /// Get statistics about the stored states
228    pub fn state_statistics(&self) -> StateStatistics {
229        let total_states = self.state_cache.len();
230        let total_memory = self.memory_footprint();
231        let compression_ratio = self.compression_ratio();
232
233        let precision_breakdown =
234            self.state_cache
235                .values()
236                .fold(HashMap::new(), |mut acc, state| {
237                    let precision_name = match state {
238                        LowPrecisionState::F16(_) => "F16",
239                        LowPrecisionState::BF16(_) => "BF16",
240                        LowPrecisionState::I8 { .. } => "I8",
241                        LowPrecisionState::I4 { .. } => "I4",
242                        LowPrecisionState::Sparse { .. } => "Sparse",
243                    };
244                    *acc.entry(precision_name.to_string()).or_insert(0) += 1;
245                    acc
246                });
247
248        StateStatistics {
249            total_states,
250            total_memory,
251            compression_ratio,
252            precision_breakdown,
253        }
254    }
255}
256
257/// Statistics about low-precision states
258#[derive(Debug, Clone)]
259pub struct StateStatistics {
260    pub total_states: usize,
261    pub total_memory: usize,
262    pub compression_ratio: f32,
263    pub precision_breakdown: HashMap<String, usize>,
264}
265
266impl fmt::Display for StateStatistics {
267    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
268        writeln!(f, "Low-Precision State Statistics:")?;
269        writeln!(f, "  Total States: {}", self.total_states)?;
270        writeln!(f, "  Total Memory: {} bytes", self.total_memory)?;
271        writeln!(f, "  Compression Ratio: {:.2}x", self.compression_ratio)?;
272        writeln!(f, "  Precision Breakdown:")?;
273        for (precision, count) in &self.precision_breakdown {
274            writeln!(f, "    {}: {} states", precision, count)?;
275        }
276        Ok(())
277    }
278}
279
280// Implement for f32 vectors (common optimizer state)
281impl LowPrecisionConvertible for Vec<f32> {
282    fn to_low_precision(&self) -> LowPrecisionState {
283        LowPrecisionState::from_f32(self, PrecisionType::F16)
284    }
285
286    fn from_low_precision(state: &LowPrecisionState) -> Result<Self, OptimizerError> {
287        Ok(state.to_f32())
288    }
289}
290
291// Implement for HashMap<String, f32> (common optimizer state)
292impl LowPrecisionConvertible for HashMap<String, f32> {
293    fn to_low_precision(&self) -> LowPrecisionState {
294        let values: Vec<f32> = self.values().copied().collect();
295        LowPrecisionState::from_f32(&values, PrecisionType::F16)
296    }
297
298    fn from_low_precision(state: &LowPrecisionState) -> Result<Self, OptimizerError> {
299        let values = state.to_f32();
300        // This is a simplified implementation - in practice, you'd need to store keys separately
301        let mut result = HashMap::new();
302        for (i, value) in values.into_iter().enumerate() {
303            result.insert(format!("param_{}", i), value);
304        }
305        Ok(result)
306    }
307}
308
309#[cfg(test)]
310mod tests {
311    use super::*;
312
313    #[test]
314    fn test_f16_conversion() {
315        let values = vec![1.0, 2.5, -3.7, 0.0, 1000.0];
316        let state = LowPrecisionState::from_f32(&values, PrecisionType::F16);
317        let recovered = state.to_f32();
318
319        // Check that values are approximately equal (f16 has limited precision)
320        for (original, recovered) in values.iter().zip(recovered.iter()) {
321            assert!(
322                (original - recovered).abs() < 0.01,
323                "Original: {}, Recovered: {}",
324                original,
325                recovered
326            );
327        }
328    }
329
330    #[test]
331    fn test_i8_conversion() {
332        let values = vec![1.0, 2.5, -3.7, 0.0, 10.0];
333        let state = LowPrecisionState::from_f32(&values, PrecisionType::I8);
334        let recovered = state.to_f32();
335
336        // I8 should have reasonable precision for small values
337        for (original, recovered) in values.iter().zip(recovered.iter()) {
338            assert!(
339                (original - recovered).abs() < 0.5,
340                "Original: {}, Recovered: {}",
341                original,
342                recovered
343            );
344        }
345    }
346
347    #[test]
348    fn test_sparse_conversion() {
349        let values = vec![0.0, 2.5, 0.0, 0.0, 10.0, 0.0];
350        let state = LowPrecisionState::from_f32(&values, PrecisionType::Sparse(1.0));
351        let recovered = state.to_f32();
352
353        // Sparse should preserve non-zero values exactly
354        assert_eq!(recovered.len(), 5); // max index + 1
355        assert_eq!(recovered[1], 2.5);
356        assert_eq!(recovered[4], 10.0);
357        for i in [0, 2, 3] {
358            assert_eq!(recovered[i], 0.0);
359        }
360    }
361
362    #[test]
363    fn test_memory_footprint() {
364        let values = vec![1.0; 1000];
365
366        let f16_state = LowPrecisionState::from_f32(&values, PrecisionType::F16);
367        let i8_state = LowPrecisionState::from_f32(&values, PrecisionType::I8);
368        let sparse_state = LowPrecisionState::from_f32(&values, PrecisionType::Sparse(2.0));
369
370        let full_size = values.len() * 4; // 4 bytes per f32
371
372        assert!(f16_state.memory_footprint() < full_size);
373        assert!(i8_state.memory_footprint() < full_size);
374        assert!(sparse_state.memory_footprint() < full_size);
375    }
376
377    #[test]
378    fn test_low_precision_optimizer() {
379        let mut optimizer =
380            LowPrecisionOptimizer::new("dummy_optimizer".to_string(), PrecisionType::F16);
381
382        let values = vec![1.0, 2.0, 3.0, 4.0];
383        optimizer.store_state("momentum".to_string(), &values);
384
385        let recovered = optimizer.load_state("momentum").unwrap();
386
387        // Check approximate equality
388        for (original, recovered) in values.iter().zip(recovered.iter()) {
389            assert!((original - recovered).abs() < 0.01);
390        }
391
392        // Check compression ratio
393        assert!(optimizer.compression_ratio() > 1.0);
394    }
395}