diff --git a/Cargo.lock b/Cargo.lock index bbbdbf06c..565003ac0 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -6264,6 +6264,7 @@ dependencies = [ "mimalloc", "ml-core", "ml-dqn", + "ml-ppo", "nalgebra 0.33.2", "ndarray", "num", @@ -6415,7 +6416,24 @@ dependencies = [ name = "ml-ppo" version = "1.0.0" dependencies = [ + "anyhow", + "approx", + "candle-core", + "candle-nn", + "candle-optimisers", + "common", "ml-core", + "ndarray", + "parking_lot 0.12.5", + "rand 0.8.5", + "rand_distr 0.4.3", + "rust_decimal", + "serde", + "serde_json", + "statrs", + "thiserror 1.0.69", + "tokio", + "tracing", ] [[package]] diff --git a/crates/ml-ppo/Cargo.toml b/crates/ml-ppo/Cargo.toml index c638990d2..ee37f782f 100644 --- a/crates/ml-ppo/Cargo.toml +++ b/crates/ml-ppo/Cargo.toml @@ -11,10 +11,43 @@ documentation.workspace = true publish.workspace = true keywords.workspace = true categories.workspace = true -description = "PPO reinforcement learning" +description = "PPO reinforcement learning for Foxhunt trading" + +[features] +default = ["cuda"] +cuda = ["candle-core/cuda", "candle-core/cudnn", "candle-nn/cuda", "candle-nn/cudnn"] [dependencies] ml-core.workspace = true +common.workspace = true + +# ML frameworks +candle-core = { git = "https://github.com/huggingface/candle", rev = "671de1db" } +candle-nn = { git = "https://github.com/huggingface/candle", rev = "671de1db" } +candle-optimisers = { git = "https://github.com/KGrewal1/optimisers" } + +# Serialization +serde = { workspace = true, features = ["derive"] } +serde_json.workspace = true + +# Core utilities +thiserror.workspace = true +anyhow.workspace = true +tracing.workspace = true +rand.workspace = true +rand_distr.workspace = true +rust_decimal.workspace = true + +# Concurrency +parking_lot = { version = "0.12", features = ["hardware-lock-elision"] } + +# Numerics +ndarray = { workspace = true, features = ["rayon"] } +statrs.workspace = true + +[dev-dependencies] +tokio = { workspace = true, features = ["test-util", "macros"] } +approx.workspace = true [lints] workspace = true diff --git a/crates/ml/src/ppo/action_masking.rs b/crates/ml-ppo/src/action_masking.rs similarity index 99% rename from crates/ml/src/ppo/action_masking.rs rename to crates/ml-ppo/src/action_masking.rs index e390cde3d..eb7a21dda 100644 --- a/crates/ml/src/ppo/action_masking.rs +++ b/crates/ml-ppo/src/action_masking.rs @@ -353,7 +353,7 @@ mod tests { #[test] fn test_45_action_masking_canonical_layout() { - use crate::common::action::{ExposureLevel, FactoredAction}; + use ml_core::action_space::{ExposureLevel, FactoredAction}; let max_position = 5.0; diff --git a/crates/ml/src/ppo/action_space.rs b/crates/ml-ppo/src/action_space.rs similarity index 97% rename from crates/ml/src/ppo/action_space.rs rename to crates/ml-ppo/src/action_space.rs index 70041f5af..91fadde1f 100644 --- a/crates/ml/src/ppo/action_space.rs +++ b/crates/ml-ppo/src/action_space.rs @@ -11,10 +11,10 @@ //! - Type-safe action handling //! - Seamless integration with existing PPO infrastructure -use crate::common::action::FactoredAction; -use crate::ppo::continuous_policy::ContinuousAction; +use ml_core::action_space::FactoredAction; +use crate::continuous_policy::ContinuousAction; use candle_core::{Tensor, Device}; -use crate::MLError; +use ml_core::MLError; use serde::{Deserialize, Serialize}; /// Action space type discriminator @@ -169,7 +169,7 @@ impl std::fmt::Display for ActionSpace { impl Default for ActionSpace { fn default() -> Self { // Default to flat exposure with market order and normal urgency - use crate::common::action::{ExposureLevel, OrderType, Urgency}; + use ml_core::action_space::{ExposureLevel, OrderType, Urgency}; ActionSpace::Discrete(FactoredAction::new( ExposureLevel::Flat, OrderType::Market, @@ -181,7 +181,7 @@ impl Default for ActionSpace { #[cfg(test)] mod tests { use super::*; - use crate::common::action::{ExposureLevel, OrderType, Urgency}; + use ml_core::action_space::{ExposureLevel, OrderType, Urgency}; use candle_core::Device; #[test] diff --git a/crates/ml/src/ppo/adaptive_entropy.rs b/crates/ml-ppo/src/adaptive_entropy.rs similarity index 99% rename from crates/ml/src/ppo/adaptive_entropy.rs rename to crates/ml-ppo/src/adaptive_entropy.rs index 6623e4138..dcf0c5d62 100644 --- a/crates/ml/src/ppo/adaptive_entropy.rs +++ b/crates/ml-ppo/src/adaptive_entropy.rs @@ -24,7 +24,7 @@ use candle_nn::{Optimizer, VarBuilder, VarMap}; use candle_optimisers::adam::{Adam, ParamsAdam}; use serde::{Deserialize, Serialize}; -use crate::MLError; +use ml_core::MLError; /// Configuration for adaptive entropy coefficient tuning. /// diff --git a/crates/ml/src/ppo/circuit_breaker.rs b/crates/ml-ppo/src/circuit_breaker.rs similarity index 62% rename from crates/ml/src/ppo/circuit_breaker.rs rename to crates/ml-ppo/src/circuit_breaker.rs index 5ff3cad71..38075a24f 100644 --- a/crates/ml/src/ppo/circuit_breaker.rs +++ b/crates/ml-ppo/src/circuit_breaker.rs @@ -1,8 +1,8 @@ //! Circuit Breaker for PPO Training //! //! Re-exports the shared circuit breaker implementation from the common module. -//! The canonical implementation lives in `crate::common::circuit_breaker`. +//! The canonical implementation lives in `ml_core::common::circuit_breaker`. -pub use crate::common::circuit_breaker::{ +pub use ml_core::common::circuit_breaker::{ CircuitBreaker, CircuitBreakerConfig, CircuitBreakerStats, CircuitState, }; diff --git a/crates/ml/src/ppo/composite_reward.rs b/crates/ml-ppo/src/composite_reward.rs similarity index 100% rename from crates/ml/src/ppo/composite_reward.rs rename to crates/ml-ppo/src/composite_reward.rs diff --git a/crates/ml/src/ppo/continuous_action_masking.rs b/crates/ml-ppo/src/continuous_action_masking.rs similarity index 99% rename from crates/ml/src/ppo/continuous_action_masking.rs rename to crates/ml-ppo/src/continuous_action_masking.rs index a8ec553ce..1bccb1b22 100644 --- a/crates/ml/src/ppo/continuous_action_masking.rs +++ b/crates/ml-ppo/src/continuous_action_masking.rs @@ -25,7 +25,7 @@ use candle_core::{Device, IndexOp, Tensor}; use serde::{Deserialize, Serialize}; -use crate::MLError; +use ml_core::MLError; /// Continuous action constraints based on current state /// diff --git a/crates/ml/src/ppo/continuous_demo.rs b/crates/ml-ppo/src/continuous_demo.rs similarity index 99% rename from crates/ml/src/ppo/continuous_demo.rs rename to crates/ml-ppo/src/continuous_demo.rs index 5b9b959be..44c532e63 100644 --- a/crates/ml/src/ppo/continuous_demo.rs +++ b/crates/ml-ppo/src/continuous_demo.rs @@ -4,7 +4,7 @@ //! for position sizing without the complex PPO training infrastructure. use super::continuous_policy::{ContinuousAction, ContinuousPolicyConfig, ContinuousPolicyNetwork}; -use crate::MLError; +use ml_core::MLError; use candle_core::{Device, Tensor}; /// Simple demo showing Gaussian policy for continuous position sizing diff --git a/crates/ml/src/ppo/continuous_policy.rs b/crates/ml-ppo/src/continuous_policy.rs similarity index 99% rename from crates/ml/src/ppo/continuous_policy.rs rename to crates/ml-ppo/src/continuous_policy.rs index 69a523bcb..b94fde0b2 100644 --- a/crates/ml/src/ppo/continuous_policy.rs +++ b/crates/ml-ppo/src/continuous_policy.rs @@ -21,9 +21,9 @@ use serde::{Deserialize, Serialize}; use statrs::distribution::{ContinuousCDF, Normal}; use tracing::{debug, warn}; -use crate::dqn::mixed_precision::training_dtype; -use crate::dqn::xavier_init::linear_xavier; -use crate::MLError; +use ml_core::mixed_precision::training_dtype; +use ml_core::xavier_init::linear_xavier; +use ml_core::MLError; /// Configuration for continuous policy network #[derive(Debug, Clone, Serialize, Deserialize)] diff --git a/crates/ml/src/ppo/continuous_ppo.rs b/crates/ml-ppo/src/continuous_ppo.rs similarity index 99% rename from crates/ml/src/ppo/continuous_ppo.rs rename to crates/ml-ppo/src/continuous_ppo.rs index f4329fe1d..9979372df 100644 --- a/crates/ml/src/ppo/continuous_ppo.rs +++ b/crates/ml-ppo/src/continuous_ppo.rs @@ -4,7 +4,7 @@ //! action spaces, using Gaussian policies for position sizing. use candle_core::{DType, Device, Tensor}; -use crate::dqn::mixed_precision::training_dtype; +use ml_core::mixed_precision::training_dtype; use candle_nn::Optimizer; // Required for Adam::new and backward_step methods use candle_optimisers::adam::Adam; use candle_optimisers::adam::ParamsAdam; @@ -15,9 +15,9 @@ use super::continuous_policy::ContinuousAction; // Keep only ContinuousAction use super::flow_policy::{FlowPolicy, FlowPolicyConfig}; use super::gae::GAEConfig; use super::ppo::ValueNetwork; -use crate::gradient_accumulation::clip_grads; -use crate::tensor_ops::TensorOps; -use crate::MLError; +use ml_core::gradient_accumulation::clip_grads; +use ml_core::tensor_ops::TensorOps; +use ml_core::MLError; /// Configuration for Continuous `PPO` #[derive(Debug, Clone, Serialize, Deserialize)] diff --git a/crates/ml/src/ppo/continuous_transaction_costs.rs b/crates/ml-ppo/src/continuous_transaction_costs.rs similarity index 99% rename from crates/ml/src/ppo/continuous_transaction_costs.rs rename to crates/ml-ppo/src/continuous_transaction_costs.rs index 96b3e996a..a6a256bf0 100644 --- a/crates/ml/src/ppo/continuous_transaction_costs.rs +++ b/crates/ml-ppo/src/continuous_transaction_costs.rs @@ -16,7 +16,7 @@ //! - LimitMaker: 5 bps (0.05%) - Passive order, maker rebate //! - IoC: 10 bps (0.10%) - Immediate-or-cancel, medium cost -pub use crate::common::action::OrderType; +pub use ml_core::action_space::OrderType; use serde::{Deserialize, Serialize}; /// Transaction cost model for continuous position sizing diff --git a/crates/ml/src/ppo/entropy_regularization.rs b/crates/ml-ppo/src/entropy_regularization.rs similarity index 99% rename from crates/ml/src/ppo/entropy_regularization.rs rename to crates/ml-ppo/src/entropy_regularization.rs index 0677e59d3..c146cbc08 100644 --- a/crates/ml/src/ppo/entropy_regularization.rs +++ b/crates/ml-ppo/src/entropy_regularization.rs @@ -23,7 +23,7 @@ use candle_core::{DType, Tensor}; -use crate::MLError; +use ml_core::MLError; /// Entropy regularizer for preventing policy collapse in PPO /// diff --git a/crates/ml/src/ppo/flow_policy/coupling_layer.rs b/crates/ml-ppo/src/flow_policy/coupling_layer.rs similarity index 99% rename from crates/ml/src/ppo/flow_policy/coupling_layer.rs rename to crates/ml-ppo/src/flow_policy/coupling_layer.rs index 1b9fbdf3a..174e4de78 100644 --- a/crates/ml/src/ppo/flow_policy/coupling_layer.rs +++ b/crates/ml-ppo/src/flow_policy/coupling_layer.rs @@ -12,8 +12,8 @@ use candle_core::{Device, Tensor}; use candle_nn::{linear, Linear, Module, VarBuilder}; -use crate::dqn::xavier_init::linear_xavier; -use crate::MLError; +use ml_core::xavier_init::linear_xavier; +use ml_core::MLError; /// Affine coupling layer with context conditioning /// diff --git a/crates/ml/src/ppo/flow_policy/flow_matching.rs b/crates/ml-ppo/src/flow_policy/flow_matching.rs similarity index 99% rename from crates/ml/src/ppo/flow_policy/flow_matching.rs rename to crates/ml-ppo/src/flow_policy/flow_matching.rs index 3085dbb25..f979163ae 100644 --- a/crates/ml/src/ppo/flow_policy/flow_matching.rs +++ b/crates/ml-ppo/src/flow_policy/flow_matching.rs @@ -12,7 +12,8 @@ use candle_core::{Device, Tensor}; -use crate::{tensor_ops::TensorOps, MLError}; +use ml_core::tensor_ops::TensorOps; +use ml_core::MLError; /// Configuration for the Flow Matching loss function. #[derive(Debug, Clone, Copy)] diff --git a/crates/ml/src/ppo/flow_policy/mod.rs b/crates/ml-ppo/src/flow_policy/mod.rs similarity index 99% rename from crates/ml/src/ppo/flow_policy/mod.rs rename to crates/ml-ppo/src/flow_policy/mod.rs index 840d6a7d0..cc81429cf 100644 --- a/crates/ml/src/ppo/flow_policy/mod.rs +++ b/crates/ml-ppo/src/flow_policy/mod.rs @@ -9,9 +9,9 @@ use rand::thread_rng; use rand_distr::{Distribution, Normal}; use serde::{Deserialize, Serialize}; -use crate::dqn::mixed_precision::training_dtype; -use crate::dqn::xavier_init::linear_xavier; -use crate::MLError; +use ml_core::mixed_precision::training_dtype; +use ml_core::xavier_init::linear_xavier; +use ml_core::MLError; /// Computes the log-determinant correction for tanh squashing. /// diff --git a/crates/ml/src/ppo/gae.rs b/crates/ml-ppo/src/gae.rs similarity index 98% rename from crates/ml/src/ppo/gae.rs rename to crates/ml-ppo/src/gae.rs index 1ec9a165f..0c6177d38 100644 --- a/crates/ml/src/ppo/gae.rs +++ b/crates/ml-ppo/src/gae.rs @@ -6,7 +6,7 @@ use serde::{Deserialize, Serialize}; use super::trajectories::Trajectory; -use crate::MLError; +use ml_core::MLError; /// Configuration for GAE computation #[derive(Debug, Clone, Serialize, Deserialize)] @@ -291,9 +291,9 @@ pub fn compute_advantages( #[cfg(test)] mod tests { use super::*; - use crate::common::action::FactoredAction; - use crate::dqn::TradingAction; - use crate::ppo::trajectories::{Trajectory, TrajectoryStep}; + use ml_core::action_space::FactoredAction; + use ml_core::trading_action::TradingAction; + use crate::trajectories::{Trajectory, TrajectoryStep}; /// Helper: build a FactoredAction from a TradingAction for test brevity. fn fa(ta: TradingAction) -> FactoredAction { diff --git a/crates/ml/src/ppo/hidden_state_manager.rs b/crates/ml-ppo/src/hidden_state_manager.rs similarity index 99% rename from crates/ml/src/ppo/hidden_state_manager.rs rename to crates/ml-ppo/src/hidden_state_manager.rs index ece3a959d..f044c24a3 100644 --- a/crates/ml/src/ppo/hidden_state_manager.rs +++ b/crates/ml-ppo/src/hidden_state_manager.rs @@ -7,8 +7,8 @@ use candle_core::{Device, Tensor}; #[cfg(test)] use candle_core::DType; use std::fmt; -use crate::MLError; -use crate::dqn::mixed_precision::training_dtype; +use ml_core::MLError; +use ml_core::mixed_precision::training_dtype; /// Manages LSTM hidden and cell states for policy and value networks pub struct HiddenStateManager { diff --git a/crates/ml-ppo/src/lib.rs b/crates/ml-ppo/src/lib.rs index 6ef21459c..0c8c2fb96 100644 --- a/crates/ml-ppo/src/lib.rs +++ b/crates/ml-ppo/src/lib.rs @@ -1 +1,60 @@ -// Modules will be moved here from ml crate +//! Proximal Policy Optimization (PPO) Implementation +//! +//! This crate provides a complete PPO implementation with: +//! - Actor-Critic architecture with separate policy and value networks +//! - Generalized Advantage Estimation (GAE) +//! - Clipped surrogate objective +//! - Trajectory collection and processing +//! - Circuit breaker for failure management +//! - Reward normalization for numerical stability +//! - Transaction costs and position limits for risk management + +// Re-export shared modules from ml-core for convenience +pub use ml_core::cuda_compat; + +pub mod adaptive_entropy; +pub mod continuous_policy; +pub mod continuous_ppo; +pub mod flow_policy; +pub mod gae; +pub mod ppo; +pub mod trajectories; +pub mod continuous_demo; +pub mod circuit_breaker; +pub mod reward_normalizer; +pub mod transaction_costs; +pub mod position_limits; +pub mod portfolio_tracker; +pub mod entropy_regularization; +pub mod action_masking; +pub mod hidden_state_manager; +pub mod lstm_networks; +pub mod action_space; +pub mod continuous_action_masking; +pub mod continuous_transaction_costs; +pub mod percentile_scaler; +pub mod reward_shaping; +pub mod symlog; +pub mod composite_reward; +pub mod trajectory_replay; + +// Re-export main components for external use +pub use continuous_policy::{ContinuousAction, ContinuousPolicyConfig, ContinuousPolicyNetwork}; +pub use continuous_ppo::{ + collect_continuous_trajectories, ContinuousPPO, ContinuousPPOConfig, ContinuousTrajectory, + ContinuousTrajectoryBatch, ContinuousTrajectoryStep, +}; +pub use flow_policy::{FlowPolicy, FlowPolicyConfig}; +pub use gae::{compute_gae, GAEConfig}; +pub use ppo::{PPOConfig, ValueNetwork, PPO}; +pub use trajectories::{Trajectory, TrajectoryBatch, TrajectorySequence, TrajectoryStep}; +pub use portfolio_tracker::PortfolioTracker; +pub use ml_core::action_space::{ExposureLevel, FactoredAction, OrderType, Urgency}; +pub use action_space::{ActionSpace, ActionType}; +pub use continuous_action_masking::{ + mask_continuous_actions, ContinuousActionConstraints, +}; +pub use continuous_transaction_costs::{ + conservative_cost_model, default_hft_cost_model, zero_cost_model, + ContinuousTransactionCosts, SlippageModel, +}; diff --git a/crates/ml/src/ppo/lstm_networks.rs b/crates/ml-ppo/src/lstm_networks.rs similarity index 98% rename from crates/ml/src/ppo/lstm_networks.rs rename to crates/ml-ppo/src/lstm_networks.rs index 3267934dd..6e51197ff 100644 --- a/crates/ml/src/ppo/lstm_networks.rs +++ b/crates/ml-ppo/src/lstm_networks.rs @@ -9,8 +9,8 @@ use candle_core::{Device, Tensor}; use candle_nn::{linear, Linear, Module, VarBuilder, VarMap, LSTM, LSTMConfig}; use candle_nn::rnn::{RNN, LSTMState}; -use crate::dqn::mixed_precision::training_dtype; -use crate::MLError; +use ml_core::mixed_precision::training_dtype; +use ml_core::MLError; /// LSTM-augmented policy network for temporal action selection /// @@ -106,7 +106,7 @@ impl LSTMPolicyNetwork { h_t: &Tensor, c_t: &Tensor, ) -> Result<(Tensor, Tensor, Tensor), MLError> { - let state = crate::dqn::mixed_precision::ensure_training_dtype(state) + let state = ml_core::mixed_precision::ensure_training_dtype(state) .map_err(|e| MLError::ModelError(e.to_string()))?; // Project input: [batch, input_dim] → [batch, hidden_dim] let x = self @@ -322,7 +322,7 @@ impl LSTMValueNetwork { h_t: &Tensor, c_t: &Tensor, ) -> Result<(Tensor, Tensor, Tensor), MLError> { - let state = crate::dqn::mixed_precision::ensure_training_dtype(state) + let state = ml_core::mixed_precision::ensure_training_dtype(state) .map_err(|e| MLError::ModelError(e.to_string()))?; // Project input: [batch, input_dim] → [batch, hidden_dim] let x = self diff --git a/crates/ml/src/ppo/percentile_scaler.rs b/crates/ml-ppo/src/percentile_scaler.rs similarity index 100% rename from crates/ml/src/ppo/percentile_scaler.rs rename to crates/ml-ppo/src/percentile_scaler.rs diff --git a/crates/ml/src/ppo/portfolio_tracker.rs b/crates/ml-ppo/src/portfolio_tracker.rs similarity index 100% rename from crates/ml/src/ppo/portfolio_tracker.rs rename to crates/ml-ppo/src/portfolio_tracker.rs diff --git a/crates/ml/src/ppo/position_limits.rs b/crates/ml-ppo/src/position_limits.rs similarity index 100% rename from crates/ml/src/ppo/position_limits.rs rename to crates/ml-ppo/src/position_limits.rs diff --git a/crates/ml/src/ppo/ppo.rs b/crates/ml-ppo/src/ppo.rs similarity index 98% rename from crates/ml/src/ppo/ppo.rs rename to crates/ml-ppo/src/ppo.rs index 66724fdd0..8841b2be2 100644 --- a/crates/ml/src/ppo/ppo.rs +++ b/crates/ml-ppo/src/ppo.rs @@ -20,20 +20,20 @@ use serde::{Deserialize, Serialize}; use std::path::PathBuf; use tracing::{debug, info, warn}; -use crate::gradient_accumulation::{accumulate_grads, check_gradients_finite, clip_grads, scale_grads}; -use crate::tensor_ops::TensorOps; +use ml_core::gradient_accumulation::{accumulate_grads, check_gradients_finite, clip_grads, scale_grads}; +use ml_core::tensor_ops::TensorOps; use super::gae::GAEConfig; use super::hidden_state_manager::HiddenStateManager; use super::lstm_networks::{LSTMPolicyNetwork, LSTMValueNetwork}; use super::trajectories::{TrajectoryBatch, TrajectoryTensors}; -use crate::dqn::circuit_breaker::{CircuitBreaker, CircuitBreakerConfig}; -use crate::dqn::mixed_precision::training_dtype; -use crate::dqn::portfolio_tracker::PortfolioTracker; -use crate::dqn::xavier_init::linear_xavier; -use crate::ppo::reward_normalizer::RewardNormalizer; -use crate::common::action::FactoredAction; -use crate::MLError; +use ml_core::common::circuit_breaker::{CircuitBreaker, CircuitBreakerConfig}; +use ml_core::mixed_precision::training_dtype; +use ml_core::portfolio_tracker::PortfolioTracker; +use ml_core::xavier_init::linear_xavier; +use crate::reward_normalizer::RewardNormalizer; +use ml_core::action_space::FactoredAction; +use ml_core::MLError; /// Actor network variants supporting both MLP and LSTM architectures /// @@ -242,7 +242,7 @@ pub struct PPOConfig { pub clip_epsilon_high: Option, /// Mixed precision configuration for BF16/FP16 forward pass on supported GPUs. /// None = FP32 only. Auto-configured based on GPU architecture at runtime. - pub mixed_precision: Option, + pub mixed_precision: Option, /// Use symlog transform for value targets (DreamerV3). Default: true. /// Compresses large returns while preserving sign. pub use_symlog: bool, @@ -299,7 +299,7 @@ pub struct PolicyNetwork { layers: Vec, device: Device, vars: VarMap, - mixed_precision: Option, + mixed_precision: Option, } impl PolicyNetwork { @@ -428,7 +428,7 @@ impl PolicyNetwork { } /// Set mixed precision config for this network - pub fn set_mixed_precision(&mut self, mp: Option) { + pub fn set_mixed_precision(&mut self, mp: Option) { self.mixed_precision = mp; } @@ -437,9 +437,9 @@ impl PolicyNetwork { pub fn forward_mixed( &self, input: &Tensor, - mixed_precision: &Option, + mixed_precision: &Option, ) -> Result { - let input = crate::dqn::mixed_precision::ensure_training_dtype(input) + let input = ml_core::mixed_precision::ensure_training_dtype(input) .map_err(|e| MLError::ModelError(e.to_string()))?; let (mut x, _use_amp) = match mixed_precision { Some(mp) if mp.enabled => { @@ -550,7 +550,7 @@ pub struct ValueNetwork { layers: Vec, device: Device, vars: VarMap, - mixed_precision: Option, + mixed_precision: Option, } impl ValueNetwork { @@ -671,7 +671,7 @@ impl ValueNetwork { } /// Set mixed precision config for this network - pub fn set_mixed_precision(&mut self, mp: Option) { + pub fn set_mixed_precision(&mut self, mp: Option) { self.mixed_precision = mp; } @@ -679,9 +679,9 @@ impl ValueNetwork { pub fn forward_mixed( &self, input: &Tensor, - mixed_precision: &Option, + mixed_precision: &Option, ) -> Result { - let input = crate::dqn::mixed_precision::ensure_training_dtype(input) + let input = ml_core::mixed_precision::ensure_training_dtype(input) .map_err(|e| MLError::ModelError(e.to_string()))?; let (mut x, _use_amp) = match mixed_precision { Some(mp) if mp.enabled => { diff --git a/crates/ml/src/ppo/reward_normalizer.rs b/crates/ml-ppo/src/reward_normalizer.rs similarity index 100% rename from crates/ml/src/ppo/reward_normalizer.rs rename to crates/ml-ppo/src/reward_normalizer.rs diff --git a/crates/ml/src/ppo/reward_shaping.rs b/crates/ml-ppo/src/reward_shaping.rs similarity index 100% rename from crates/ml/src/ppo/reward_shaping.rs rename to crates/ml-ppo/src/reward_shaping.rs diff --git a/crates/ml/src/ppo/symlog.rs b/crates/ml-ppo/src/symlog.rs similarity index 99% rename from crates/ml/src/ppo/symlog.rs rename to crates/ml-ppo/src/symlog.rs index 47ebf05df..071ae3cf9 100644 --- a/crates/ml/src/ppo/symlog.rs +++ b/crates/ml-ppo/src/symlog.rs @@ -15,7 +15,7 @@ use candle_core::Tensor; -use crate::MLError; +use ml_core::MLError; /// Symlog transform: `sign(x) * ln(|x| + 1)` /// diff --git a/crates/ml/src/ppo/trajectories.rs b/crates/ml-ppo/src/trajectories.rs similarity index 99% rename from crates/ml/src/ppo/trajectories.rs rename to crates/ml-ppo/src/trajectories.rs index 52137cde3..f77cfb188 100644 --- a/crates/ml/src/ppo/trajectories.rs +++ b/crates/ml-ppo/src/trajectories.rs @@ -6,8 +6,8 @@ use candle_core::Tensor; use serde::{Deserialize, Serialize}; -use crate::common::action::FactoredAction; -use crate::MLError; +use ml_core::action_space::FactoredAction; +use ml_core::MLError; /// Single step trajectory data #[derive(Debug, Clone, Serialize, Deserialize)] @@ -610,8 +610,8 @@ where #[cfg(test)] mod tests { use super::*; - use crate::common::action::{ExposureLevel, OrderType, Urgency}; - use crate::dqn::TradingAction; + use ml_core::action_space::{ExposureLevel, OrderType, Urgency}; + use ml_core::trading_action::TradingAction; /// Helper: build a FactoredAction from a TradingAction for test brevity. fn fa(ta: TradingAction) -> FactoredAction { diff --git a/crates/ml/src/ppo/trajectory_replay.rs b/crates/ml-ppo/src/trajectory_replay.rs similarity index 100% rename from crates/ml/src/ppo/trajectory_replay.rs rename to crates/ml-ppo/src/trajectory_replay.rs diff --git a/crates/ml/src/ppo/transaction_costs.rs b/crates/ml-ppo/src/transaction_costs.rs similarity index 98% rename from crates/ml/src/ppo/transaction_costs.rs rename to crates/ml-ppo/src/transaction_costs.rs index 8289c54b6..4ec26a2b2 100644 --- a/crates/ml/src/ppo/transaction_costs.rs +++ b/crates/ml-ppo/src/transaction_costs.rs @@ -8,7 +8,7 @@ //! - **LimitMaker**: 0.05% (0.0005) - Passive order, maker rebate //! - **IoC**: 0.10% (0.0010) - Immediate-or-cancel, medium cost -pub use crate::common::action::OrderType; +pub use ml_core::action_space::OrderType; /// Calculate total transaction cost for a trade /// diff --git a/crates/ml/Cargo.toml b/crates/ml/Cargo.toml index c7e93b941..6993947af 100644 --- a/crates/ml/Cargo.toml +++ b/crates/ml/Cargo.toml @@ -70,6 +70,7 @@ colored = "2.1" # Terminal color output for evaluation reports # Internal workspace crates ml-core.workspace = true ml-dqn.workspace = true +ml-ppo.workspace = true config.workspace = true common = { workspace = true, features = ["questdb"] } risk = { path = "../risk" } diff --git a/crates/ml/src/ppo/mod.rs b/crates/ml/src/ppo/mod.rs index 2b7a70a9d..ad50b70f8 100644 --- a/crates/ml/src/ppo/mod.rs +++ b/crates/ml/src/ppo/mod.rs @@ -1,62 +1,15 @@ //! Proximal Policy Optimization (PPO) Implementation //! -//! This module provides a complete PPO implementation with: -//! - Actor-Critic architecture with separate policy and value networks -//! - Generalized Advantage Estimation (GAE) -//! - Clipped surrogate objective -//! - Trajectory collection and processing -//! - Real mathematical operations using candle-core v0.9.1 -//! - Circuit breaker for failure management -//! - Reward normalization for numerical stability -//! - Transaction costs and position limits for risk management +//! This module re-exports the `ml-ppo` crate and keeps bridge modules that depend +//! on types from both `ml-ppo` and the parent `ml` crate (e.g., `UnifiedTrainable`, +//! `PPOTrainer`). -pub mod adaptive_entropy; -pub mod continuous_policy; -pub mod continuous_ppo; -pub mod flow_policy; -pub mod gae; -pub mod ppo; -pub mod trajectories; -// pub mod continuous_example; // Module file not found -pub mod continuous_demo; +// Re-export everything from the ml-ppo crate +pub use ml_ppo::*; + +// Bridge modules that depend on ml-internal types (UnifiedTrainable, PPOTrainer) pub mod trainable_adapter; -pub mod circuit_breaker; -pub mod reward_normalizer; -pub mod transaction_costs; -pub mod position_limits; -pub mod portfolio_tracker; -pub mod entropy_regularization; -pub mod action_masking; pub mod stress_testing; -pub mod hidden_state_manager; -pub mod lstm_networks; -pub mod action_space; -pub mod continuous_action_masking; -pub mod continuous_transaction_costs; -pub mod percentile_scaler; -pub mod reward_shaping; -pub mod symlog; -pub mod composite_reward; -pub mod trajectory_replay; -// Re-export main components for external use -pub use continuous_policy::{ContinuousAction, ContinuousPolicyConfig, ContinuousPolicyNetwork}; -pub use continuous_ppo::{ - collect_continuous_trajectories, ContinuousPPO, ContinuousPPOConfig, ContinuousTrajectory, - ContinuousTrajectoryBatch, ContinuousTrajectoryStep, -}; -pub use flow_policy::{FlowPolicy, FlowPolicyConfig}; -pub use gae::{compute_gae, GAEConfig}; -pub use ppo::{PPOConfig, ValueNetwork, PPO}; +// Re-export bridge types pub use trainable_adapter::{train_batch, UnifiedPPO as UnifiedTrainablePPO}; -pub use trajectories::{Trajectory, TrajectoryBatch, TrajectorySequence, TrajectoryStep}; -pub use portfolio_tracker::PortfolioTracker; -pub use crate::common::action::{ExposureLevel, FactoredAction, OrderType, Urgency}; -pub use action_space::{ActionSpace, ActionType}; -pub use continuous_action_masking::{ - mask_continuous_actions, ContinuousActionConstraints, -}; -pub use continuous_transaction_costs::{ - conservative_cost_model, default_hft_cost_model, zero_cost_model, - ContinuousTransactionCosts, SlippageModel, -};