use crate::error::{InferenceError, InferenceResult};
use kizzasi_core::HiddenState;
use scirs2_core::ndarray::Array2;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CompressionMethod {
None,
Quantize8Bit,
Quantize4Bit,
Sparse,
QuantizedSparse,
}
#[derive(Debug, Clone)]
pub struct CompressedState {
method: CompressionMethod,
data: Vec<u8>,
shape: Vec<usize>,
scale: f32,
zero_point: i32,
sparse_indices: Option<Vec<usize>>,
}
impl CompressedState {
pub fn compression_ratio(&self) -> f32 {
let original_size = self.shape.iter().product::<usize>() * std::mem::size_of::<f32>();
let compressed_size = self.data.len()
+ self
.sparse_indices
.as_ref()
.map(|v| v.len() * std::mem::size_of::<usize>())
.unwrap_or(0);
original_size as f32 / compressed_size as f32
}
pub fn method(&self) -> CompressionMethod {
self.method
}
}
pub struct StateCompressor {
method: CompressionMethod,
sparsity_threshold: f32,
}
impl StateCompressor {
pub fn new(method: CompressionMethod) -> Self {
Self {
method,
sparsity_threshold: 1e-4,
}
}
pub fn with_sparsity_threshold(mut self, threshold: f32) -> Self {
self.sparsity_threshold = threshold;
self
}
pub fn compress(&self, state: &HiddenState) -> InferenceResult<CompressedState> {
match self.method {
CompressionMethod::None => self.compress_none(state),
CompressionMethod::Quantize8Bit => self.compress_quantize_8bit(state),
CompressionMethod::Quantize4Bit => self.compress_quantize_4bit(state),
CompressionMethod::Sparse => self.compress_sparse(state),
CompressionMethod::QuantizedSparse => self.compress_quantized_sparse(state),
}
}
pub fn decompress(&self, compressed: &CompressedState) -> InferenceResult<HiddenState> {
match compressed.method {
CompressionMethod::None => self.decompress_none(compressed),
CompressionMethod::Quantize8Bit => self.decompress_quantize_8bit(compressed),
CompressionMethod::Quantize4Bit => self.decompress_quantize_4bit(compressed),
CompressionMethod::Sparse => self.decompress_sparse(compressed),
CompressionMethod::QuantizedSparse => self.decompress_quantized_sparse(compressed),
}
}
fn compress_none(&self, state: &HiddenState) -> InferenceResult<CompressedState> {
let data_vec: Vec<f32> = state.state().iter().copied().collect();
let data_bytes: Vec<u8> = data_vec.iter().flat_map(|&f| f.to_le_bytes()).collect();
let shape_vec: Vec<usize> = state.state().shape().to_vec();
Ok(CompressedState {
method: CompressionMethod::None,
data: data_bytes,
shape: shape_vec,
scale: 1.0,
zero_point: 0,
sparse_indices: None,
})
}
fn decompress_none(&self, compressed: &CompressedState) -> InferenceResult<HiddenState> {
let floats: Vec<f32> = compressed
.data
.chunks_exact(4)
.map(|chunk| f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]))
.collect();
let data = Array2::from_shape_vec((compressed.shape[0], compressed.shape[1]), floats)
.map_err(|e| InferenceError::ForwardError(e.to_string()))?;
let mut hidden = HiddenState::new(compressed.shape[0], compressed.shape[1]);
hidden.update(data);
Ok(hidden)
}
fn compress_quantize_8bit(&self, state: &HiddenState) -> InferenceResult<CompressedState> {
let min_val = state.state().iter().copied().fold(f32::INFINITY, f32::min);
let max_val = state
.state()
.iter()
.copied()
.fold(f32::NEG_INFINITY, f32::max);
let scale = (max_val - min_val) / 255.0;
let zero_point = (-min_val / scale).round() as i32;
let quantized: Vec<u8> = state
.state()
.iter()
.map(|&v| {
let scaled = (v / scale + zero_point as f32).round();
scaled.clamp(0.0, 255.0) as u8
})
.collect();
let shape_vec: Vec<usize> = state.state().shape().to_vec();
Ok(CompressedState {
method: CompressionMethod::Quantize8Bit,
data: quantized,
shape: shape_vec,
scale,
zero_point,
sparse_indices: None,
})
}
fn decompress_quantize_8bit(
&self,
compressed: &CompressedState,
) -> InferenceResult<HiddenState> {
let dequantized: Vec<f32> = compressed
.data
.iter()
.map(|&q| (q as f32 - compressed.zero_point as f32) * compressed.scale)
.collect();
let data = Array2::from_shape_vec((compressed.shape[0], compressed.shape[1]), dequantized)
.map_err(|e| InferenceError::ForwardError(e.to_string()))?;
let mut hidden = HiddenState::new(compressed.shape[0], compressed.shape[1]);
hidden.update(data);
Ok(hidden)
}
fn compress_quantize_4bit(&self, state: &HiddenState) -> InferenceResult<CompressedState> {
let min_val = state.state().iter().copied().fold(f32::INFINITY, f32::min);
let max_val = state
.state()
.iter()
.copied()
.fold(f32::NEG_INFINITY, f32::max);
let scale = (max_val - min_val) / 15.0;
let zero_point = (-min_val / scale).round() as i32;
let mut quantized = Vec::new();
let mut iter = state.state().iter();
while let Some(&v1) = iter.next() {
let q1 = ((v1 / scale + zero_point as f32).round().clamp(0.0, 15.0) as u8) & 0x0F;
let q2 = if let Some(&v2) = iter.next() {
((v2 / scale + zero_point as f32).round().clamp(0.0, 15.0) as u8) & 0x0F
} else {
0
};
quantized.push((q1 << 4) | q2);
}
let shape_vec: Vec<usize> = state.state().shape().to_vec();
Ok(CompressedState {
method: CompressionMethod::Quantize4Bit,
data: quantized,
shape: shape_vec,
scale,
zero_point,
sparse_indices: None,
})
}
fn decompress_quantize_4bit(
&self,
compressed: &CompressedState,
) -> InferenceResult<HiddenState> {
let total_elements = compressed.shape.iter().product();
let mut dequantized = Vec::with_capacity(total_elements);
for &byte in &compressed.data {
let q1 = (byte >> 4) & 0x0F;
let q2 = byte & 0x0F;
dequantized.push((q1 as f32 - compressed.zero_point as f32) * compressed.scale);
if dequantized.len() < total_elements {
dequantized.push((q2 as f32 - compressed.zero_point as f32) * compressed.scale);
}
}
let data = Array2::from_shape_vec((compressed.shape[0], compressed.shape[1]), dequantized)
.map_err(|e| InferenceError::ForwardError(e.to_string()))?;
let mut hidden = HiddenState::new(compressed.shape[0], compressed.shape[1]);
hidden.update(data);
Ok(hidden)
}
fn compress_sparse(&self, state: &HiddenState) -> InferenceResult<CompressedState> {
let mut values = Vec::new();
let mut indices = Vec::new();
for (i, &v) in state.state().iter().enumerate() {
if v.abs() > self.sparsity_threshold {
values.push(v);
indices.push(i);
}
}
let data_bytes: Vec<u8> = values.iter().flat_map(|&f| f.to_le_bytes()).collect();
let shape_vec: Vec<usize> = state.state().shape().to_vec();
Ok(CompressedState {
method: CompressionMethod::Sparse,
data: data_bytes,
shape: shape_vec,
scale: 1.0,
zero_point: 0,
sparse_indices: Some(indices),
})
}
fn decompress_sparse(&self, compressed: &CompressedState) -> InferenceResult<HiddenState> {
let values: Vec<f32> = compressed
.data
.chunks_exact(4)
.map(|chunk| f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]))
.collect();
let indices = compressed
.sparse_indices
.as_ref()
.ok_or(InferenceError::ForwardError(
"Missing sparse indices".to_string(),
))?;
let total_elements: usize = compressed.shape.iter().product();
let mut dense = vec![0.0f32; total_elements];
for (&idx, &val) in indices.iter().zip(values.iter()) {
if idx < dense.len() {
dense[idx] = val;
}
}
let data = Array2::from_shape_vec((compressed.shape[0], compressed.shape[1]), dense)
.map_err(|e| InferenceError::ForwardError(e.to_string()))?;
let mut hidden = HiddenState::new(compressed.shape[0], compressed.shape[1]);
hidden.update(data);
Ok(hidden)
}
fn compress_quantized_sparse(&self, state: &HiddenState) -> InferenceResult<CompressedState> {
let mut values = Vec::new();
let mut indices = Vec::new();
for (i, &v) in state.state().iter().enumerate() {
if v.abs() > self.sparsity_threshold {
values.push(v);
indices.push(i);
}
}
let shape_vec: Vec<usize> = state.state().shape().to_vec();
if values.is_empty() {
return Ok(CompressedState {
method: CompressionMethod::QuantizedSparse,
data: Vec::new(),
shape: shape_vec,
scale: 1.0,
zero_point: 0,
sparse_indices: Some(indices),
});
}
let min_val = values.iter().copied().fold(f32::INFINITY, f32::min);
let max_val = values.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let scale = (max_val - min_val) / 255.0;
let zero_point = (-min_val / scale).round() as i32;
let quantized: Vec<u8> = values
.iter()
.map(|&v| {
let scaled = (v / scale + zero_point as f32).round();
scaled.clamp(0.0, 255.0) as u8
})
.collect();
Ok(CompressedState {
method: CompressionMethod::QuantizedSparse,
data: quantized,
shape: shape_vec,
scale,
zero_point,
sparse_indices: Some(indices),
})
}
fn decompress_quantized_sparse(
&self,
compressed: &CompressedState,
) -> InferenceResult<HiddenState> {
let indices = compressed
.sparse_indices
.as_ref()
.ok_or(InferenceError::ForwardError(
"Missing sparse indices".to_string(),
))?;
let total_elements: usize = compressed.shape.iter().product();
let mut dense = vec![0.0f32; total_elements];
if !compressed.data.is_empty() {
for (&idx, &q) in indices.iter().zip(compressed.data.iter()) {
if idx < dense.len() {
dense[idx] = (q as f32 - compressed.zero_point as f32) * compressed.scale;
}
}
}
let data = Array2::from_shape_vec((compressed.shape[0], compressed.shape[1]), dense)
.map_err(|e| InferenceError::ForwardError(e.to_string()))?;
let mut hidden = HiddenState::new(compressed.shape[0], compressed.shape[1]);
hidden.update(data);
Ok(hidden)
}
}