1use crate::OptimizerError;
2use std::collections::HashMap;
3use std::fmt;
4
5pub trait LowPrecisionConvertible: Clone + fmt::Debug {
7 fn to_low_precision(&self) -> LowPrecisionState;
9
10 fn from_low_precision(state: &LowPrecisionState) -> Result<Self, OptimizerError>;
12}
13
14#[derive(Clone, Debug)]
16pub enum LowPrecisionState {
17 F16(Vec<half::f16>),
19 BF16(Vec<half::bf16>),
21 I8 { values: Vec<i8>, scale: f32 },
23 I4 { values: Vec<u8>, scale: f32 }, Sparse {
27 indices: Vec<usize>,
28 values: Vec<f32>,
29 },
30}
31
32impl LowPrecisionState {
33 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, LowPrecisionState::I4 { values, .. } => (values.len() + 1) / 2 + 4, LowPrecisionState::Sparse { indices, values } => {
41 indices.len() * 8 + values.len() * 4 }
43 }
44 }
45
46 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 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#[derive(Clone, Debug)]
140pub enum PrecisionType {
141 F16,
143 BF16,
145 I8,
147 I4,
149 Sparse(f32),
151}
152
153pub struct LowPrecisionOptimizer<T> {
155 inner: T,
156 precision: PrecisionType,
157 state_cache: HashMap<String, LowPrecisionState>,
158}
159
160impl<T> LowPrecisionOptimizer<T> {
161 pub fn new(inner: T, precision: PrecisionType) -> Self {
163 Self {
164 inner,
165 precision,
166 state_cache: HashMap::new(),
167 }
168 }
169
170 pub fn precision(&self) -> &PrecisionType {
172 &self.precision
173 }
174
175 pub fn inner(&self) -> &T {
177 &self.inner
178 }
179
180 pub fn inner_mut(&mut self) -> &mut T {
182 &mut self.inner
183 }
184
185 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 pub fn load_state(&self, key: &str) -> Option<Vec<f32>> {
193 self.state_cache.get(key).map(|state| state.to_f32())
194 }
195
196 pub fn memory_footprint(&self) -> usize {
198 self.state_cache
199 .values()
200 .map(|state| state.memory_footprint())
201 .sum()
202 }
203
204 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; let compressed_size = self.memory_footprint();
218
219 full_precision_size as f32 / compressed_size as f32
220 }
221
222 pub fn clear_cache(&mut self) {
224 self.state_cache.clear();
225 }
226
227 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#[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
280impl 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
291impl 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 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 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 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 assert_eq!(recovered.len(), 5); 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; 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 for (original, recovered) in values.iter().zip(recovered.iter()) {
389 assert!((original - recovered).abs() < 0.01);
390 }
391
392 assert!(optimizer.compression_ratio() > 1.0);
394 }
395}