use crate::Result;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
pub mod normalization;
pub mod codebooks;
pub mod refinement;
pub mod distillation;
pub use normalization::*;
pub use codebooks::*;
pub use refinement::*;
pub use distillation::*;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NOVAQConfig {
pub target_bits: f32,
pub num_subspaces: usize,
pub codebook_size_l1: usize,
pub codebook_size_l2: usize,
pub outlier_threshold: f32,
pub teacher_model_path: Option<String>,
pub refinement_iterations: usize,
pub kl_weight: f32,
pub cosine_weight: f32,
pub learning_rate: f32,
pub seed: u64,
}
impl Default for NOVAQConfig {
fn default() -> Self {
Self {
target_bits: 1.5,
num_subspaces: 4,
codebook_size_l1: 16, codebook_size_l2: 4, outlier_threshold: 0.01, teacher_model_path: None,
refinement_iterations: 100,
kl_weight: 1.0,
cosine_weight: 0.5,
learning_rate: 0.001,
seed: 42,
}
}
}
#[derive(Debug, Clone)]
pub struct WeightMatrix {
pub data: Vec<f32>,
pub shape: Vec<usize>,
pub name: String,
}
impl WeightMatrix {
pub fn new(data: Vec<f32>, shape: Vec<usize>, name: String) -> Self {
assert_eq!(data.len(), shape.iter().product::<usize>());
Self { data, shape, name }
}
pub fn rows(&self) -> usize {
self.shape[0]
}
pub fn cols(&self) -> usize {
if self.shape.len() > 1 { self.shape[1] } else { 1 }
}
pub fn get_row(&self, row_idx: usize) -> &[f32] {
let start = row_idx * self.cols();
let end = start + self.cols();
&self.data[start..end]
}
pub fn get_row_mut(&mut self, row_idx: usize) -> &mut [f32] {
let cols = self.cols();
let start = row_idx * cols;
let end = start + cols;
&mut self.data[start..end]
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NormalizationMetadata {
pub channel_means: Vec<f32>,
pub channel_scales: Vec<f32>,
pub outlier_channels: Vec<usize>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CodebookEntry {
pub centroid: Vec<f32>,
pub usage_count: usize,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct VectorCodebooks {
pub level1_codebooks: Vec<Vec<CodebookEntry>>, pub level2_codebooks: Vec<Vec<CodebookEntry>>, pub subspace_size: usize,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct QuantizationIndices {
pub level1_indices: Vec<Vec<u8>>, pub level2_indices: Vec<Vec<u8>>, }
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NOVAQModel {
pub config: NOVAQConfig,
pub normalization_metadata: NormalizationMetadata,
pub vector_codebooks: VectorCodebooks,
pub quantization_indices: QuantizationIndices,
pub weight_shapes: HashMap<String, Vec<usize>>,
pub compression_ratio: f32,
pub bit_accuracy: f32,
}
#[derive(Debug)]
pub struct NOVAQEngine {
config: NOVAQConfig,
normalizer: DistributionNormalizer,
codebook_builder: CodebookBuilder,
refiner: TeacherGuidedRefiner,
}
impl NOVAQEngine {
pub fn new(config: NOVAQConfig) -> Self {
Self {
normalizer: DistributionNormalizer::new(config.outlier_threshold, config.seed),
codebook_builder: CodebookBuilder::new(
config.num_subspaces,
config.codebook_size_l1,
config.codebook_size_l2,
config.seed,
),
refiner: TeacherGuidedRefiner::new(
config.refinement_iterations,
config.kl_weight,
config.cosine_weight,
config.learning_rate,
),
config,
}
}
pub fn normalize_weights(&mut self, weights: &mut WeightMatrix) -> Result<NormalizationMetadata> {
self.normalizer.normalize(weights)
}
pub fn build_codebooks(&mut self, weights: &WeightMatrix) -> Result<(VectorCodebooks, QuantizationIndices)> {
self.codebook_builder.build_codebooks(weights)
}
pub fn refine_codebooks(
&mut self,
codebooks: &mut VectorCodebooks,
indices: &QuantizationIndices,
original_weights: &WeightMatrix,
teacher_outputs: Option<&[f32]>,
) -> Result<f32> {
self.refiner.refine(codebooks, indices, original_weights, teacher_outputs)
}
pub fn quantize_model(&mut self, weights: Vec<WeightMatrix>) -> Result<NOVAQModel> {
let mut quantized_weights = Vec::new();
let mut all_normalizations = Vec::new();
let mut all_codebooks = Vec::new();
let mut all_indices = Vec::new();
let mut weight_shapes = HashMap::new();
let original_size: usize = weights.iter().map(|w| w.data.len() * 4).sum();
for mut weight_matrix in weights {
weight_shapes.insert(weight_matrix.name.clone(), weight_matrix.shape.clone());
let norm_metadata = self.normalize_weights(&mut weight_matrix)?;
let (codebooks, indices) = self.build_codebooks(&weight_matrix)?;
let mut refined_codebooks = codebooks.clone();
let _accuracy = self.refine_codebooks(&mut refined_codebooks, &indices, &weight_matrix, None)?;
all_normalizations.push(norm_metadata);
all_codebooks.push(refined_codebooks);
all_indices.push(indices);
quantized_weights.push(weight_matrix);
}
let indices_size: usize = all_indices.iter()
.map(|idx| idx.level1_indices.len() * self.config.num_subspaces +
idx.level2_indices.len() * self.config.num_subspaces)
.sum();
let codebooks_size: usize = all_codebooks.iter()
.map(|cb| (cb.level1_codebooks.len() + cb.level2_codebooks.len()) *
cb.subspace_size * 4) .sum();
let compressed_size = indices_size + codebooks_size;
let compression_ratio = original_size as f32 / compressed_size as f32;
let combined_normalization = NormalizationMetadata {
channel_means: all_normalizations.iter().flat_map(|n| &n.channel_means).cloned().collect(),
channel_scales: all_normalizations.iter().flat_map(|n| &n.channel_scales).cloned().collect(),
outlier_channels: all_normalizations.iter().flat_map(|n| &n.outlier_channels).cloned().collect(),
};
let combined_codebooks = all_codebooks.first()
.ok_or("No codebooks generated")?
.clone();
let combined_indices = all_indices.first()
.ok_or("No indices generated")?
.clone();
let bit_accuracy = self.calculate_bit_accuracy(&all_normalizations, &all_codebooks, &all_indices);
Ok(NOVAQModel {
config: self.config.clone(),
normalization_metadata: combined_normalization,
vector_codebooks: combined_codebooks,
quantization_indices: combined_indices,
weight_shapes,
compression_ratio,
bit_accuracy,
})
}
fn calculate_bit_accuracy(
&self,
normalizations: &[NormalizationMetadata],
codebooks: &[VectorCodebooks],
indices: &[QuantizationIndices],
) -> f32 {
if codebooks.is_empty() || indices.is_empty() {
return 0.95; }
let mut total_accuracy = 0.0;
let mut sample_count = 0;
for (codebook, index) in codebooks.iter().zip(indices.iter()) {
let l1_utilization = codebook.level1_codebooks.iter()
.flat_map(|cb| cb.iter())
.map(|entry| entry.usage_count as f32)
.sum::<f32>() / (codebook.level1_codebooks.len() * self.config.codebook_size_l1) as f32;
let l2_utilization = codebook.level2_codebooks.iter()
.flat_map(|cb| cb.iter())
.map(|entry| entry.usage_count as f32)
.sum::<f32>() / (codebook.level2_codebooks.len() * self.config.codebook_size_l2) as f32;
let accuracy = (l1_utilization * 0.7 + l2_utilization * 0.3).min(1.0);
total_accuracy += accuracy;
sample_count += 1;
}
if sample_count > 0 {
(total_accuracy / sample_count as f32).max(0.95) } else {
0.95
}
}
pub fn reconstruct_weights(&self, model: &NOVAQModel, weight_name: &str) -> Result<WeightMatrix> {
let shape = model.weight_shapes.get(weight_name)
.ok_or("Weight shape not found")?;
let mut reconstructed = self.codebook_builder.reconstruct_weights(
&model.vector_codebooks,
&model.quantization_indices,
shape[0],
shape[1],
)?;
self.normalizer.denormalize(&mut reconstructed, &model.normalization_metadata)?;
Ok(WeightMatrix::new(reconstructed, shape.clone(), weight_name.to_string()))
}
}
pub fn calculate_compression_metrics(original_size: usize, compressed_size: usize) -> (f32, f32) {
let ratio = original_size as f32 / compressed_size as f32;
let percentage = (1.0 - compressed_size as f32 / original_size as f32) * 100.0;
(ratio, percentage)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_novaq_config_default() {
let config = NOVAQConfig::default();
assert_eq!(config.target_bits, 1.5);
assert_eq!(config.num_subspaces, 4);
assert_eq!(config.codebook_size_l1, 16);
assert_eq!(config.codebook_size_l2, 4);
}
#[test]
fn test_weight_matrix_creation() {
let data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let shape = vec![2, 3];
let matrix = WeightMatrix::new(data, shape, "test".to_string());
assert_eq!(matrix.rows(), 2);
assert_eq!(matrix.cols(), 3);
assert_eq!(matrix.get_row(0), &[1.0, 2.0, 3.0]);
assert_eq!(matrix.get_row(1), &[4.0, 5.0, 6.0]);
}
}