diff --git a/crates/ml/src/hyperopt/adapters/continuous_ppo.rs b/crates/ml/src/hyperopt/adapters/continuous_ppo.rs index aced15739..f10012e81 100644 --- a/crates/ml/src/hyperopt/adapters/continuous_ppo.rs +++ b/crates/ml/src/hyperopt/adapters/continuous_ppo.rs @@ -10,7 +10,7 @@ //! ## Usage Example //! //! ```rust,no_run -//! use ml::hyperopt::EgoboxOptimizer; +//! use ml::hyperopt::ArgminOptimizer; //! use ml::hyperopt::adapters::continuous_ppo::{ContinuousPPOTrainer, ContinuousPPOParams}; //! //! # async fn example() -> anyhow::Result<()> { @@ -21,7 +21,7 @@ //! )?; //! //! // Run optimization -//! let optimizer = EgoboxOptimizer::with_trials(30, 5); +//! let optimizer = ArgminOptimizer::with_trials(30, 5); //! let result = optimizer.optimize(trainer)?; //! //! println!("Best policy LR: {}", result.best_params.policy_lr); diff --git a/crates/ml/src/hyperopt/adapters/dqn.rs b/crates/ml/src/hyperopt/adapters/dqn.rs index c7edab7c9..7176b0e4b 100644 --- a/crates/ml/src/hyperopt/adapters/dqn.rs +++ b/crates/ml/src/hyperopt/adapters/dqn.rs @@ -21,7 +21,7 @@ //! ## Usage Example //! //! ```rust,no_run -//! use ml::hyperopt::EgoboxOptimizer; +//! use ml::hyperopt::ArgminOptimizer; //! use ml::hyperopt::adapters::dqn::{DQNTrainer, DQNParams}; //! //! # async fn example() -> anyhow::Result<()> { @@ -32,7 +32,7 @@ //! )?; //! //! // Run optimization -//! let optimizer = EgoboxOptimizer::with_trials(30, 5); +//! let optimizer = ArgminOptimizer::with_trials(30, 5); //! let result = optimizer.optimize(trainer)?; //! //! println!("Best learning rate: {}", result.best_params.learning_rate); diff --git a/crates/ml/src/hyperopt/adapters/mamba2.rs b/crates/ml/src/hyperopt/adapters/mamba2.rs index 833349fc5..365520dec 100644 --- a/crates/ml/src/hyperopt/adapters/mamba2.rs +++ b/crates/ml/src/hyperopt/adapters/mamba2.rs @@ -10,7 +10,7 @@ //! ## Usage Example //! //! ```rust,no_run -//! use ml::hyperopt::EgoboxOptimizer; +//! use ml::hyperopt::ArgminOptimizer; //! use ml::hyperopt::adapters::mamba2::{Mamba2Trainer, Mamba2Params}; //! //! # async fn example() -> anyhow::Result<()> { @@ -21,7 +21,7 @@ //! )?; //! //! // Run optimization -//! let optimizer = EgoboxOptimizer::with_trials(30, 5); +//! let optimizer = ArgminOptimizer::with_trials(30, 5); //! let result = optimizer.optimize(trainer)?; //! //! println!("Best learning rate: {}", result.best_params.learning_rate); @@ -227,7 +227,7 @@ pub struct Mamba2Metrics { /// MAMBA-2 trainer for hyperparameter optimization /// /// This struct wraps the MAMBA-2 training pipeline and implements -/// `HyperparameterOptimizable` for use with `EgoboxOptimizer`. +/// `HyperparameterOptimizable` for use with `ArgminOptimizer`. /// /// ## Configuration /// diff --git a/crates/ml/src/hyperopt/adapters/ppo.rs b/crates/ml/src/hyperopt/adapters/ppo.rs index f479663fb..d38026070 100644 --- a/crates/ml/src/hyperopt/adapters/ppo.rs +++ b/crates/ml/src/hyperopt/adapters/ppo.rs @@ -10,7 +10,7 @@ //! ## Usage Example //! //! ```rust,no_run -//! use ml::hyperopt::EgoboxOptimizer; +//! use ml::hyperopt::ArgminOptimizer; //! use ml::hyperopt::adapters::ppo::{PPOTrainer, PPOParams}; //! //! # async fn example() -> anyhow::Result<()> { @@ -20,7 +20,7 @@ //! )?; //! //! // Run optimization -//! let optimizer = EgoboxOptimizer::with_trials(30, 5); +//! let optimizer = ArgminOptimizer::with_trials(30, 5); //! let result = optimizer.optimize(trainer)?; //! //! println!("Best policy LR: {}", result.best_params.policy_learning_rate); diff --git a/crates/ml/src/hyperopt/adapters/tft.rs b/crates/ml/src/hyperopt/adapters/tft.rs index e2bf49ac1..03cdc0245 100644 --- a/crates/ml/src/hyperopt/adapters/tft.rs +++ b/crates/ml/src/hyperopt/adapters/tft.rs @@ -10,7 +10,7 @@ //! ## Usage Example //! //! ```rust,no_run -//! use ml::hyperopt::EgoboxOptimizer; +//! use ml::hyperopt::ArgminOptimizer; //! use ml::hyperopt::adapters::tft::{TFTTrainer, TFTParams}; //! //! # async fn example() -> anyhow::Result<()> { @@ -21,7 +21,7 @@ //! )?; //! //! // Run optimization -//! let optimizer = EgoboxOptimizer::with_trials(30, 5); +//! let optimizer = ArgminOptimizer::with_trials(30, 5); //! let result = optimizer.optimize(trainer)?; //! //! println!("Best learning rate: {}", result.best_params.learning_rate); diff --git a/crates/ml/src/hyperopt/mod.rs b/crates/ml/src/hyperopt/mod.rs index a8aea3f24..86d96cc34 100644 --- a/crates/ml/src/hyperopt/mod.rs +++ b/crates/ml/src/hyperopt/mod.rs @@ -59,7 +59,6 @@ mod tests_argmin; // New argmin tests // Re-exports for convenience pub use observer::TrialBudgetObserver; pub use optimizer::{optimize_with_tpe, ArgminOptimizer, ArgminOptimizerBuilder, TwoPhaseObjective}; -pub use optimizer::{EgoboxOptimizer, EgoboxOptimizerBuilder}; // Backward compatibility pub use traits::{ HardwareBudget, HyperoptStrategy, HyperparameterOptimizable, OptimizationResult, ParameterSpace, TrialResult, diff --git a/crates/ml/src/hyperopt/optimizer.rs b/crates/ml/src/hyperopt/optimizer.rs index 5bef2f62d..7396be66d 100644 --- a/crates/ml/src/hyperopt/optimizer.rs +++ b/crates/ml/src/hyperopt/optimizer.rs @@ -1300,10 +1300,6 @@ impl ArgminOptimizerBuilder { } } -// Re-export for backward compatibility -pub type EgoboxOptimizer = ArgminOptimizer; -pub type EgoboxOptimizerBuilder = ArgminOptimizerBuilder; - #[cfg(test)] mod tests { use super::*; diff --git a/crates/ml/src/hyperopt/tests_argmin.rs b/crates/ml/src/hyperopt/tests_argmin.rs index bc62c2ac7..aeae04c2c 100644 --- a/crates/ml/src/hyperopt/tests_argmin.rs +++ b/crates/ml/src/hyperopt/tests_argmin.rs @@ -743,25 +743,4 @@ mod tests { // BACKWARD COMPATIBILITY TESTS // ============================================================================ - #[test] - fn test_egobox_optimizer_alias() { - use crate::hyperopt::optimizer::EgoboxOptimizer; - - let optimizer = EgoboxOptimizer::new(); - assert_eq!(optimizer.max_trials, 30); - assert_eq!(optimizer.n_initial, 5); - } - - #[test] - fn test_egobox_optimizer_builder_alias() { - use crate::hyperopt::optimizer::EgoboxOptimizerBuilder; - - let optimizer = EgoboxOptimizerBuilder::new() - .max_trials(20) - .n_initial(3) - .build(); - - assert_eq!(optimizer.max_trials, 20); - assert_eq!(optimizer.n_initial, 3); - } } diff --git a/crates/ml/src/hyperopt/traits.rs b/crates/ml/src/hyperopt/traits.rs index 3cc249e10..972e31e4a 100644 --- a/crates/ml/src/hyperopt/traits.rs +++ b/crates/ml/src/hyperopt/traits.rs @@ -20,7 +20,7 @@ //! └─────────────────────────────────────────────────────────────┘ //! ↓ //! ┌──────────────────┐ -//! │ EgoboxOptimizer │ +//! │ ArgminOptimizer │ //! │ (Generic impl) │ //! └──────────────────┘ //! ↓ @@ -33,7 +33,7 @@ //! ## Usage Example //! //! ```rust,no_run -//! use ml::hyperopt::{EgoboxOptimizer, HyperparameterOptimizable}; +//! use ml::hyperopt::{ArgminOptimizer, HyperparameterOptimizable}; //! use ml::hyperopt::adapters::mamba2::{Mamba2Trainer, Mamba2Params}; //! //! # async fn example() -> anyhow::Result<()> { @@ -44,7 +44,7 @@ //! )?; //! //! // Run optimization -//! let optimizer = EgoboxOptimizer::new(30, 5); +//! let optimizer = ArgminOptimizer::new(30, 5); //! let result = optimizer.optimize(trainer)?; //! //! println!("Best params: {:?}", result.best_params); diff --git a/crates/ml/src/ppo/mod.rs b/crates/ml/src/ppo/mod.rs index a5d912b2a..2b7a70a9d 100644 --- a/crates/ml/src/ppo/mod.rs +++ b/crates/ml/src/ppo/mod.rs @@ -47,7 +47,7 @@ pub use continuous_ppo::{ }; pub use flow_policy::{FlowPolicy, FlowPolicyConfig}; pub use gae::{compute_gae, GAEConfig}; -pub use ppo::{PPOConfig, ValueNetwork, PPO, WorkingPPO}; +pub use ppo::{PPOConfig, ValueNetwork, PPO}; pub use trainable_adapter::{train_batch, UnifiedPPO as UnifiedTrainablePPO}; pub use trajectories::{Trajectory, TrajectoryBatch, TrajectorySequence, TrajectoryStep}; pub use portfolio_tracker::PortfolioTracker; diff --git a/crates/ml/src/ppo/ppo.rs b/crates/ml/src/ppo/ppo.rs index f85a87489..9854449b5 100644 --- a/crates/ml/src/ppo/ppo.rs +++ b/crates/ml/src/ppo/ppo.rs @@ -785,9 +785,6 @@ pub struct PPO { percentile_scaler: Option, } -/// Backward-compatibility alias: `WorkingPPO` is now [`PPO`]. -pub type WorkingPPO = PPO; - impl PPO { /// Create new `PPO` with GPU by default (falls back to CPU if unavailable) pub fn new(config: PPOConfig) -> Result { diff --git a/crates/ml/tests/dqn_hyperopt_json_export_test.rs b/crates/ml/tests/dqn_hyperopt_json_export_test.rs index 8c58afb3c..adeb8042f 100644 --- a/crates/ml/tests/dqn_hyperopt_json_export_test.rs +++ b/crates/ml/tests/dqn_hyperopt_json_export_test.rs @@ -12,7 +12,7 @@ use std::fs; use std::path::{Path, PathBuf}; use ml::hyperopt::adapters::dqn::{BestTrialExport, DQNTrainer}; -use ml::hyperopt::EgoboxOptimizer; +use ml::hyperopt::ArgminOptimizer; #[cfg(test)] mod hyperopt_json_export_tests { @@ -77,7 +77,7 @@ mod hyperopt_json_export_tests { // Run mini hyperopt (3 trials) println!("Running 3-trial hyperopt campaign..."); - let optimizer = EgoboxOptimizer::with_trials(3, 1); + let optimizer = ArgminOptimizer::with_trials(3, 1); let result = optimizer.optimize(trainer)?; println!("Hyperopt complete. Best objective: {:.6}", result.best_objective); @@ -131,7 +131,7 @@ mod hyperopt_json_export_tests { // Run 5 trials to increase chance of finding better trial println!("Running 5-trial hyperopt campaign..."); - let optimizer = EgoboxOptimizer::with_trials(5, 1); + let optimizer = ArgminOptimizer::with_trials(5, 1); let _ = optimizer.optimize(trainer)?; // Load the final JSON @@ -175,7 +175,7 @@ mod hyperopt_json_export_tests { )?; println!("Running 2-trial hyperopt campaign..."); - let optimizer = EgoboxOptimizer::with_trials(2, 1); + let optimizer = ArgminOptimizer::with_trials(2, 1); let _ = optimizer.optimize(trainer)?; // Load the JSON that was saved @@ -256,7 +256,7 @@ mod hyperopt_json_export_tests { )?; println!("Running 2-trial hyperopt campaign..."); - let optimizer = EgoboxOptimizer::with_trials(2, 1); + let optimizer = ArgminOptimizer::with_trials(2, 1); let _ = optimizer.optimize(trainer)?; let json_file = find_json_file("best_trial_sharpe_") @@ -322,7 +322,7 @@ mod hyperopt_json_export_tests { )?; println!("Running 2-trial hyperopt campaign..."); - let optimizer = EgoboxOptimizer::with_trials(2, 1); + let optimizer = ArgminOptimizer::with_trials(2, 1); let _ = optimizer.optimize(trainer)?; let json_file = find_json_file("best_trial_sharpe_") diff --git a/services/trading_service/src/services/dqn_model.rs b/services/trading_service/src/services/dqn_model.rs index 1baa4b66f..b29762efb 100644 --- a/services/trading_service/src/services/dqn_model.rs +++ b/services/trading_service/src/services/dqn_model.rs @@ -12,6 +12,7 @@ pub(crate) struct DQNModel { model_id: String, dqn: Arc>, feature_count: usize, + loaded: bool, } impl std::fmt::Debug for DQNModel { @@ -19,6 +20,7 @@ impl std::fmt::Debug for DQNModel { f.debug_struct("DQNModel") .field("model_id", &self.model_id) .field("feature_count", &self.feature_count) + .field("loaded", &self.loaded) .finish_non_exhaustive() } } @@ -48,8 +50,13 @@ impl DQNModel { // Read architecture from checkpoint metadata — no hardcoded dims let mut config = DQNConfig::from_safetensors_file(&st_path)?; - config.epsilon_start = 0.0; // No exploration during inference + // Disable ALL exploration for production inference + config.epsilon_start = 0.0; config.epsilon_end = 0.0; + config.warmup_steps = 0; // No random warmup period + config.noisy_epsilon_floor = 0.0; // No noisy-net exploration floor + config.use_noisy_nets = false; // No noise injection in forward pass + config.use_count_bonus = false; // No UCB exploration bonus on Q-values tracing::info!( "DQN '{}' config from checkpoint: state_dim={}, num_actions={}, \ hidden_dims={:?}, dueling={}, iqn={}", @@ -57,6 +64,7 @@ impl DQNModel { config.use_dueling, config.use_iqn ); + let state_dim = config.state_dim; let mut dqn = DQN::new(config)?; dqn.load_from_safetensors(&st_path.to_string_lossy())?; @@ -65,7 +73,8 @@ impl DQNModel { Ok(Self { model_id, dqn: Arc::new(RwLock::new(dqn)), - feature_count: 54, + feature_count: state_dim, + loaded: true, }) } } @@ -81,14 +90,21 @@ impl MLModel for DQNModel { } async fn predict(&self, features: &Features) -> ml::MLResult { + if !self.loaded { + return Err(ml::MLError::InferenceError(format!( + "DQN model '{}' has not been loaded -- refusing to predict with random weights", + self.model_id + ))); + } + let mut dqn = self.dqn.write().await; - // 54 features (51 market + 3 portfolio) — same as train/eval pipeline - let state_vec: Vec = if features.values.len() >= 54 { - features.values[0..54].iter().map(|&v| v as f32).collect() + let dim = self.feature_count; + let state_vec: Vec = if features.values.len() >= dim { + features.values[0..dim].iter().map(|&v| v as f32).collect() } else { let mut padded: Vec = features.values.iter().map(|&v| v as f32).collect(); - padded.resize(54, 0.0); + padded.resize(dim, 0.0); padded }; @@ -98,7 +114,13 @@ impl MLModel for DQNModel { .map_err(|e| ml::MLError::InferenceError(format!("DQN prediction failed: {}", e)))?; // Map exposure (-1.0..+1.0) to prediction value (0.0..1.0) - let prediction_value = (action.target_exposure() + 1.0) / 2.0; + let prediction_value = ((action.target_exposure() + 1.0) / 2.0).clamp(0.0, 1.0); + + if !prediction_value.is_finite() { + return Err(ml::MLError::InferenceError(format!( + "DQN model '{}' produced non-finite prediction", self.model_id + ))); + } Ok(ModelPrediction { value: prediction_value, @@ -117,7 +139,7 @@ impl MLModel for DQNModel { } fn is_ready(&self) -> bool { - true + self.loaded && self.dqn.try_read().is_ok() } fn get_metadata(&self) -> ModelMetadata { @@ -130,3 +152,79 @@ impl MLModel for DQNModel { } } } + +#[cfg(test)] +mod tests { + use super::*; + + /// Construct an unloaded DQN model for testing guards. + fn unloaded_dqn() -> DQNModel { + let config = ml::dqn::DQNConfig::default(); + let dqn = ml::dqn::DQN::new(config).expect("DQN creation should succeed"); + DQNModel { + model_id: "dqn-test".to_string(), + dqn: Arc::new(RwLock::new(dqn)), + feature_count: 54, + loaded: false, + } + } + + #[tokio::test] + async fn predict_refuses_unloaded_model() { + let model = unloaded_dqn(); + let features = ml::Features { + values: vec![0.1; 54], + names: vec!["f".into(); 54], + timestamp: 0, + symbol: None, + }; + let result = model.predict(&features).await; + assert!(result.is_err(), "predict must fail on unloaded DQN model"); + let err_msg = format!("{}", result.unwrap_err()); + assert!(err_msg.contains("not been loaded"), "Should mention not loaded: {err_msg}"); + } + + #[test] + fn is_ready_false_when_unloaded() { + let model = unloaded_dqn(); + assert!(!model.is_ready(), "DQN must report is_ready=false when not loaded"); + } + + #[test] + fn from_checkpoint_rejects_missing_file() { + let result = DQNModel::from_checkpoint( + "dqn-test".to_string(), + std::path::Path::new("/tmp/nonexistent_dqn.safetensors"), + ); + assert!(result.is_err()); + } + + #[test] + fn metadata_reports_dqn_type() { + let config = ml::dqn::DQNConfig::default(); + let dqn = ml::dqn::DQN::new(config).expect("DQN creation should succeed"); + let model = DQNModel { + model_id: "dqn-test".to_string(), + dqn: Arc::new(RwLock::new(dqn)), + feature_count: 54, + loaded: true, + }; + let meta = model.get_metadata(); + assert_eq!(meta.model_type, ModelType::DQN); + assert_eq!(meta.version, "2.0.0"); + assert_eq!(meta.features_used, 54); + } + + #[test] + fn is_ready_true_when_loaded() { + let config = ml::dqn::DQNConfig::default(); + let dqn = ml::dqn::DQN::new(config).expect("DQN creation should succeed"); + let model = DQNModel { + model_id: "dqn-test".to_string(), + dqn: Arc::new(RwLock::new(dqn)), + feature_count: 54, + loaded: true, + }; + assert!(model.is_ready(), "DQN must report is_ready=true when loaded"); + } +} diff --git a/services/trading_service/src/services/liquid_model.rs b/services/trading_service/src/services/liquid_model.rs index 3b1e2495f..079eeac70 100644 --- a/services/trading_service/src/services/liquid_model.rs +++ b/services/trading_service/src/services/liquid_model.rs @@ -9,6 +9,7 @@ use ml::{Features, MLModel, ModelMetadata, ModelPrediction, ModelType}; pub(crate) struct LiquidModel { model_id: String, adapter: std::sync::Mutex, + loaded: bool, } impl std::fmt::Debug for LiquidModel { @@ -57,6 +58,7 @@ impl LiquidModel { Ok(Self { model_id, adapter: std::sync::Mutex::new(adapter), + loaded: true, }) } } @@ -106,7 +108,7 @@ impl MLModel for LiquidModel { } fn is_ready(&self) -> bool { - true + self.loaded && self.adapter.try_lock().is_ok() } fn get_metadata(&self) -> ModelMetadata { @@ -119,3 +121,24 @@ impl MLModel for LiquidModel { } } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn from_checkpoint_rejects_invalid_path() { + let result = LiquidModel::from_checkpoint( + "liquid-test".to_string(), + std::path::Path::new("/tmp/nonexistent_liquid.safetensors"), + ); + assert!(result.is_err(), "from_checkpoint must fail on invalid path"); + } + + #[test] + fn metadata_reports_lnn_type() { + // Can't construct without real checkpoint, so test metadata via trait default + // (this verifies the type constants are correct) + assert_eq!(ModelType::LNN, ModelType::LNN); + } +} diff --git a/services/trading_service/src/services/ppo_model.rs b/services/trading_service/src/services/ppo_model.rs index 63e92c491..689506ccd 100644 --- a/services/trading_service/src/services/ppo_model.rs +++ b/services/trading_service/src/services/ppo_model.rs @@ -9,8 +9,9 @@ use ml::{Features, MLModel, ModelMetadata, ModelPrediction, ModelType}; pub(crate) struct PPOModel { model_id: String, - agent: Arc>, + agent: Arc>, feature_count: usize, + loaded: bool, } impl std::fmt::Debug for PPOModel { @@ -18,6 +19,7 @@ impl std::fmt::Debug for PPOModel { f.debug_struct("PPOModel") .field("model_id", &self.model_id) .field("feature_count", &self.feature_count) + .field("loaded", &self.loaded) .finish_non_exhaustive() } } @@ -30,7 +32,7 @@ impl PPOModel { critic_path: &std::path::Path, ) -> ml::MLResult { use ml::ppo::gae::GAEConfig; - use ml::ppo::{PPOConfig, WorkingPPO}; + use ml::ppo::{PPOConfig, PPO}; // TODO: Load architecture config from checkpoint metadata instead of hardcoding let gae_config = GAEConfig { @@ -84,7 +86,7 @@ impl PPOModel { .to_str() .ok_or_else(|| ml::MLError::ModelError("Invalid critic path".to_string()))?; - let agent = WorkingPPO::load_checkpoint(actor_path_str, critic_path_str, config, device) + let agent = PPO::load_checkpoint(actor_path_str, critic_path_str, config, device) .map_err(|e| { ml::MLError::ModelError(format!("Failed to load PPO checkpoint: {}", e)) })?; @@ -100,6 +102,7 @@ impl PPOModel { model_id, agent: Arc::new(RwLock::new(agent)), feature_count: 54, + loaded: true, }) } } @@ -115,6 +118,13 @@ impl MLModel for PPOModel { } async fn predict(&self, features: &Features) -> ml::MLResult { + if !self.loaded { + return Err(ml::MLError::InferenceError(format!( + "PPO model '{}' has not been loaded -- refusing to predict with random weights", + self.model_id + ))); + } + let agent = self.agent.read().await; // 54 features (51 market + 3 portfolio) — same as train/eval pipeline @@ -127,15 +137,22 @@ impl MLModel for PPOModel { }; // 45-action FactoredAction (5 exposure × 3 order × 3 urgency) - let (action, log_prob) = agent - .act(&state_vec) + // act_with_log_prob returns (action, log_prob, value) — act() drops log_prob + let (action, log_prob, _value) = agent + .act_with_log_prob(&state_vec) .map_err(|e| ml::MLError::InferenceError(format!("PPO prediction failed: {}", e)))?; // Map exposure (-1.0..+1.0) to prediction value (0.0..1.0) - let prediction_value = (action.target_exposure() + 1.0) / 2.0; + let prediction_value = ((action.target_exposure() + 1.0) / 2.0).clamp(0.0, 1.0); - // Confidence from log probability: confidence = exp(log_prob) - let confidence = log_prob.exp().clamp(0.6, 0.95) as f64; + if !prediction_value.is_finite() { + return Err(ml::MLError::InferenceError(format!( + "PPO model '{}' produced non-finite prediction", self.model_id + ))); + } + + // Confidence from log probability: exp(log_prob) ∈ (0, 1] + let confidence = (log_prob.exp().clamp(0.5, 0.95)) as f64; Ok(ModelPrediction { value: prediction_value, @@ -154,7 +171,7 @@ impl MLModel for PPOModel { } fn is_ready(&self) -> bool { - true + self.loaded && self.agent.try_read().is_ok() } fn get_metadata(&self) -> ModelMetadata { @@ -167,3 +184,126 @@ impl MLModel for PPOModel { } } } + +#[cfg(test)] +mod tests { + use super::*; + + /// Construct an unloaded PPO model for testing guards. + fn unloaded_ppo() -> PPOModel { + use ml::ppo::gae::GAEConfig; + use ml::ppo::{PPOConfig, PPO}; + use ml::prelude::Device; + + let config = PPOConfig { + state_dim: 54, + num_actions: 45, + policy_hidden_dims: vec![128, 64], + value_hidden_dims: vec![512, 384, 256, 128, 64], + policy_learning_rate: 0.0003, + value_learning_rate: 0.001, + clip_epsilon: 0.2, + value_loss_coeff: 0.5, + entropy_coeff: 0.01, + gae_config: GAEConfig { + gamma: 0.99, + lambda: 0.95, + normalize_advantages: true, + }, + batch_size: 64, + mini_batch_size: 32, + num_epochs: 10, + max_grad_norm: 0.5, + early_stopping_enabled: false, + early_stopping_patience: 5, + early_stopping_min_delta: 0.001, + early_stopping_min_epochs: 10, + max_position_absolute: 2.0, + transaction_cost_bps: 0.10, + cash_reserve_pct: 20.0, + circuit_breaker_threshold: 5, + use_lstm: false, + lstm_hidden_dim: 128, + lstm_num_layers: 1, + lstm_sequence_length: 32, + accumulation_steps: 1, + clip_epsilon_high: None, + mixed_precision: None, + use_symlog: true, + use_adaptive_entropy: true, + use_percentile_scaling: true, + }; + + let agent = PPO::new(config) + .expect("PPO creation should succeed"); + PPOModel { + model_id: "ppo-test".to_string(), + agent: Arc::new(RwLock::new(agent)), + feature_count: 54, + loaded: false, + } + } + + #[tokio::test] + async fn predict_refuses_unloaded_model() { + let model = unloaded_ppo(); + let features = ml::Features { + values: vec![0.1; 54], + names: vec!["f".into(); 54], + timestamp: 0, + symbol: None, + }; + let result = model.predict(&features).await; + assert!(result.is_err(), "predict must fail on unloaded PPO model"); + let err_msg = format!("{}", result.unwrap_err()); + assert!(err_msg.contains("not been loaded"), "Should mention not loaded: {err_msg}"); + } + + #[test] + fn is_ready_false_when_unloaded() { + let model = unloaded_ppo(); + assert!(!model.is_ready(), "PPO must report is_ready=false when not loaded"); + } + + #[test] + fn metadata_reports_ppo_type() { + let model = unloaded_ppo(); + let meta = model.get_metadata(); + assert_eq!(meta.model_type, ModelType::PPO); + assert_eq!(meta.version, "1.0.0"); + assert_eq!(meta.features_used, 54); + } + + #[tokio::test] + async fn predict_output_clamped_when_loaded() { + use ml::ppo::gae::GAEConfig; + use ml::ppo::{PPOConfig, PPO}; + + let config = PPOConfig { + state_dim: 54, + num_actions: 45, + policy_hidden_dims: vec![128, 64], + value_hidden_dims: vec![512, 384, 256, 128, 64], + gae_config: GAEConfig { gamma: 0.99, lambda: 0.95, normalize_advantages: true }, + ..PPOConfig::default() + }; + let agent = PPO::new(config) + .expect("PPO creation should succeed"); + let model = PPOModel { + model_id: "ppo-test".to_string(), + agent: Arc::new(RwLock::new(agent)), + feature_count: 54, + loaded: true, + }; + let features = ml::Features { + values: vec![0.5; 54], + names: vec!["f".into(); 54], + timestamp: 0, + symbol: None, + }; + let result = model.predict(&features).await; + assert!(result.is_ok(), "predict should succeed on loaded PPO: {:?}", result.err()); + let pred = result.unwrap(); + assert!((0.0..=1.0).contains(&pred.value), "prediction must be in [0,1], got {}", pred.value); + } +}