diff --git a/services/trading_service/src/services/enhanced_ml.rs b/services/trading_service/src/services/enhanced_ml.rs index c8c3867fe..fb452ef9e 100644 --- a/services/trading_service/src/services/enhanced_ml.rs +++ b/services/trading_service/src/services/enhanced_ml.rs @@ -1689,9 +1689,9 @@ struct RealTFTModel { impl RealTFTModel { /// Create new TFT model from checkpoint /// - /// Validates that the checkpoint file exists before creating the model. - /// Returns an error if the checkpoint is missing -- never sets is_trained=true - /// on random weights. + /// Loads actual trained weights from a safetensors checkpoint file. + /// Returns an error if the checkpoint does not exist or cannot be loaded, + /// preventing production trading on random (untrained) weights. pub fn from_checkpoint( model_id: String, checkpoint_path: &std::path::Path, @@ -1707,10 +1707,18 @@ impl RealTFTModel { } info!( - "Initializing TFT model (checkpoint: {})", + "Loading TFT model from checkpoint: {}", checkpoint_path.display() ); + // Verify checkpoint file exists before creating model + if !checkpoint_path.exists() { + return Err(ml::MLError::CheckpointError(format!( + "TFT checkpoint file not found: {}. Cannot use random weights in production.", + checkpoint_path.display() + ))); + } + // TFT configuration matching Wave 160 training // input_dim must equal num_static + num_known + num_unknown = 5 + 10 + 16 = 31 let config = TFTConfig { @@ -1735,18 +1743,35 @@ impl RealTFTModel { target_throughput_pps: 100_000, }; - // Create TFT model -- checkpoint weight loading not yet implemented in ml crate, - // so we do NOT set is_trained=true. The model will refuse to produce predictions - // until real weight loading is added. - let tft = TemporalFusionTransformer::new(config.clone()) + // Create TFT model structure + let mut tft = TemporalFusionTransformer::new(config.clone()) .map_err(|e| ml::MLError::ModelError(format!("Failed to create TFT: {}", e)))?; - // NOTE: is_trained is NOT set to true. The checkpoint file exists but the ml crate - // does not yet support loading TFT weights from safetensors. Callers must check - // is_ready() before using this model. - warn!( - "TFT model {} created with architecture only -- safetensors weight loading not yet implemented in ml crate. \ - Checkpoint exists at {} but weights are random. is_trained=false.", + // Load trained weights from safetensors checkpoint into the VarMap + let checkpoint_str = checkpoint_path + .to_str() + .ok_or_else(|| ml::MLError::ModelError("Invalid TFT checkpoint path encoding".to_string()))?; + + // VarMap::load populates the existing VarMap with weights from the safetensors file. + // Arc::get_mut is safe here because we just created tft and hold the only reference. + let varmap_mut = Arc::get_mut(tft.varmap_mut()).ok_or_else(|| { + ml::MLError::ModelError( + "Cannot load TFT checkpoint: VarMap has multiple references".to_string(), + ) + })?; + varmap_mut.load(checkpoint_str).map_err(|e| { + ml::MLError::CheckpointError(format!( + "Failed to load TFT weights from {}: {}", + checkpoint_path.display(), + e + )) + })?; + + // Only mark as trained after weights are successfully loaded + tft.is_trained = true; + + info!( + "Loaded TFT model {} from checkpoint: {}", model_id, checkpoint_path.display() ); @@ -1849,9 +1874,9 @@ struct RealMamba2Model { impl RealMamba2Model { /// Create new Mamba2 model from checkpoint /// - /// Validates that the checkpoint file exists before creating the model. - /// Returns an error if the checkpoint is missing -- never sets is_trained=true - /// on random weights. + /// Loads actual trained weights from a safetensors checkpoint file. + /// Returns an error if the checkpoint does not exist or cannot be loaded, + /// preventing production trading on random (untrained) weights. pub fn from_checkpoint( model_id: String, checkpoint_path: &std::path::Path, @@ -1867,25 +1892,51 @@ impl RealMamba2Model { } info!( - "Initializing Mamba2 model (checkpoint: {})", + "Loading Mamba2 model from checkpoint: {}", checkpoint_path.display() ); + // Verify checkpoint file exists before creating model + if !checkpoint_path.exists() { + return Err(ml::MLError::CheckpointError(format!( + "Mamba2 checkpoint file not found: {}. Cannot use random weights in production.", + checkpoint_path.display() + ))); + } + let config = Mamba2Config::default(); use ml::prelude::Device; let device = Device::Cpu; - // Create Mamba2 model -- checkpoint weight loading not yet implemented, - // so we do NOT set is_trained=true. - let mamba2 = Mamba2SSM::new(config.clone(), &device) + // Create Mamba2 model structure + let mut mamba2 = Mamba2SSM::new(config.clone(), &device) .map_err(|e| ml::MLError::ModelError(format!("Failed to create Mamba2: {}", e)))?; - // NOTE: is_trained is NOT set to true. The checkpoint file exists but the ml crate - // does not yet support loading Mamba2 weights from safetensors. Callers must check - // is_ready() before using this model. - warn!( - "Mamba2 model {} created with architecture only -- safetensors weight loading not yet implemented. \ - Checkpoint exists at {} but weights are random. is_trained=false.", + // Load trained weights from safetensors checkpoint into the VarMap. + let checkpoint_str = checkpoint_path + .to_str() + .ok_or_else(|| ml::MLError::ModelError("Invalid Mamba2 checkpoint path encoding".to_string()))?; + + // VarMap::load populates the existing VarMap with weights from the safetensors file. + // Arc::get_mut is safe here because we just created mamba2 and hold the only reference. + let varmap_mut = Arc::get_mut(&mut mamba2.varmap).ok_or_else(|| { + ml::MLError::ModelError( + "Cannot load Mamba2 checkpoint: VarMap has multiple references".to_string(), + ) + })?; + varmap_mut.load(checkpoint_str).map_err(|e| { + ml::MLError::CheckpointError(format!( + "Failed to load Mamba2 weights from {}: {}", + checkpoint_path.display(), + e + )) + })?; + + // Only mark as trained after weights are successfully loaded + mamba2.is_trained = true; + + info!( + "Loaded Mamba2 model {} from checkpoint: {}", model_id, checkpoint_path.display() );