Quantize a FP32 vector to ternary trits using BitNet absmean scaling. Returns `(trits, scale)` where `scale = mean(|v_i|)`.
(v: &[f32])
| 106 | /// |
| 107 | /// Returns `(trits, scale)` where `scale = mean(|v_i|)`. |
| 108 | pub fn quantize(v: &[f32]) -> (Vec<i8>, f32) { |
| 109 | if v.is_empty() { |
| 110 | return (Vec::new(), 0.0); |
| 111 | } |
| 112 | let scale: f32 = v.iter().map(|x| x.abs()).sum::<f32>() / v.len() as f32; |
| 113 | let trits = if scale == 0.0 { |
| 114 | vec![0i8; v.len()] |
| 115 | } else { |
| 116 | v.iter() |
| 117 | .map(|&x| (x / scale).round().clamp(-1.0, 1.0) as i8) |
| 118 | .collect() |
| 119 | }; |
| 120 | (trits, scale) |
| 121 | } |
| 122 | |
| 123 | #[cfg(test)] |
| 124 | mod tests { |