trajectories.rs: MiniBatch + TrajectorySequence CPU structs deleted. create_mini_batches() → create_mini_batch_ranges() (range indices only) to_sequences() → to_sequence_ranges() (range + length only) DtoD sub-batch extraction via CudaTrajectoryTensors::sub_batch() continuous_ppo.rs: ContinuousMiniBatch CPU struct deleted. create_mini_batches() → create_mini_batch_ranges() returning (usize,usize) trajectory_tensors.rs: sub_batch(start, end, stream) using memcpy_dtod_async for zero-CPU mini-batch slicing ppo.rs: compute_losses() uploads ONCE, iterates ranges with sub_batch() 168/168 tests pass. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
76 lines
3.1 KiB
Rust
76 lines
3.1 KiB
Rust
//! 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
|
|
|
|
#![allow(clippy::module_name_repetitions)]
|
|
#![allow(clippy::integer_division)]
|
|
#![allow(clippy::shadow_reuse, clippy::shadow_same, clippy::shadow_unrelated)] // Tensor ops: let x = x.relu() is idiomatic
|
|
#![allow(clippy::non_ascii_literal)] // Math symbols in ML documentation and error messages
|
|
#![allow(clippy::partial_pub_fields)] // ML config structs: some fields are pub API, some internal
|
|
#![allow(clippy::same_name_method)] // Intentional: inherent methods shadow trait defaults for ML-specific behavior
|
|
#![allow(clippy::indexing_slicing)] // Tensor/matrix indexing with bounds guaranteed by construction
|
|
#![allow(clippy::similar_names)] // ML naming: min_val/max_val, state/states are conventional
|
|
#![allow(unsafe_code)] // Required for CUDA kernel launches and cuBLAS FFI
|
|
|
|
// Re-export shared modules from ml-core for convenience
|
|
pub use ml_core::cuda_compile;
|
|
|
|
pub mod adaptive_entropy;
|
|
pub mod continuous_policy;
|
|
pub mod continuous_ppo;
|
|
pub mod cuda_nn;
|
|
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, PolicyNetwork, PPO};
|
|
pub use cuda_nn::{
|
|
GpuContext, CudaPolicyNetwork, CudaValueNetwork, CudaTrajectoryTensors,
|
|
CudaLinear, CudaLSTM, CudaAdam,
|
|
};
|
|
pub use trajectories::{Trajectory, TrajectoryBatch, TrajectoryStep, MiniBatchRange, SequenceRange};
|
|
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_host as mask_continuous_actions, ContinuousActionConstraints,
|
|
};
|
|
pub use continuous_transaction_costs::{
|
|
conservative_cost_model, default_hft_cost_model, zero_cost_model,
|
|
ContinuousTransactionCosts, SlippageModel,
|
|
};
|