//! Model loading and management for the inference engine. use std::collections::HashMap; use std::sync::Arc; use std::time::{Duration, Instant}; use rtx_tensor::{DType, Tensor}; use tokio::sync::RwLock; use tracing::info; use crate::quantization::{QuantizationConfig, QuantizationScheme}; use crate::{InferenceError, InferenceResult}; use super::types::{ModelConfig, ModelInfo, ModelMetrics, OptimizationStats}; /// Internal model representation #[derive(Clone)] pub struct LoadedModel { pub info: ModelInfo, pub weights: HashMap, pub config: ModelConfig, pub optimization_stats: OptimizationStats, pub metrics: Arc>, } /// Model management operations pub struct ModelManager; impl ModelManager { /// Validate model weights against configuration pub fn validate_model_weights( weights: &HashMap, config: &ModelConfig, ) -> InferenceResult<()> { if weights.is_empty() { return Err(InferenceError::invalid_request( "Model weights cannot be empty", )); } // Validate embedding layer if let Some(embedding) = weights.get("embedding.weight") { let shape = embedding.shape(); if shape.dims().len() != 2 || shape.dims()[0] != config.vocab_size || shape.dims()[1] != config.hidden_size { return Err(InferenceError::invalid_request(format!( "Invalid embedding shape: expected [{}, {}], got {:?}", config.vocab_size, config.hidden_size, shape.dims() ))); } } Ok(()) } /// Optimize model graph with various optimizations pub fn optimize_model_graph( weights: &mut HashMap, _config: &ModelConfig, ) -> InferenceResult { let original_count = weights.len(); let mut optimizations_applied = 0; let mut memory_saved = 0; // Constant folding optimization let mut to_remove = vec![]; for (name, tensor) in weights.iter() { if name.contains("bias") && tensor.shape().dims().iter().product::() == 0 { to_remove.push(name.clone()); memory_saved += 1024; optimizations_applied += 1; } } for name in to_remove { weights.remove(&name); } let nodes_eliminated = original_count - weights.len(); let estimated_speedup = if optimizations_applied > 0 { 1.0 + (optimizations_applied as f64 * 0.1) } else { 1.0 }; Ok(OptimizationStats { optimizations_applied, nodes_eliminated, memory_saved, estimated_speedup, }) } /// Calculate memory usage for model weights #[must_use] pub fn calculate_memory_usage(weights: &HashMap) -> usize { weights .values() .map(|tensor| { tensor.shape().dims().iter().product::() * match tensor.dtype() { DType::F32 => 4, DType::F16 => 2, DType::BF16 => 2, _ => 4, } }) .sum() } /// Calculate parameter count for model weights #[must_use] pub fn calculate_parameter_count(weights: &HashMap) -> usize { weights .values() .map(|tensor| tensor.shape().dims().iter().product::()) .sum() } /// Create a new loaded model #[must_use] pub fn create_loaded_model( name: &str, version: &str, config: &ModelConfig, weights: HashMap, optimization_stats: OptimizationStats, ) -> LoadedModel { let parameter_count = Self::calculate_parameter_count(&weights); let memory_usage = Self::calculate_memory_usage(&weights); let model_info = ModelInfo { name: name.to_string(), version: version.to_string(), config: config.clone(), layer_count: weights.len(), parameter_count, memory_usage, original_memory_usage: memory_usage, is_quantized: false, is_optimized: optimization_stats.optimizations_applied > 0, load_time: Instant::now(), }; LoadedModel { info: model_info, weights, config: config.clone(), optimization_stats: optimization_stats.clone(), metrics: Arc::new(RwLock::new(ModelMetrics { inference_count: 0, total_inference_time: Duration::ZERO, average_tokens_per_request: 0.0, cache_hit_rate: 0.0, optimization_speedup: optimization_stats.estimated_speedup, })), } } /// Apply quantization to a model pub fn apply_quantization( model: &mut LoadedModel, config: &QuantizationConfig, ) -> InferenceResult<()> { let original_memory = model.info.memory_usage; let quantization_factor = match config.scheme { QuantizationScheme::INT8 => 0.25, QuantizationScheme::INT4 => 0.125, QuantizationScheme::FP8E4M3 | QuantizationScheme::FP8E5M2 => 0.25, }; model.info.memory_usage = (original_memory as f64 * quantization_factor) as usize; model.info.is_quantized = true; info!( "Quantization applied, memory reduced from {} to {} bytes", original_memory, model.info.memory_usage ); Ok(()) } }