diff --git a/crates/ml/src/dqn/dqn.rs b/crates/ml/src/dqn/dqn.rs index 28b659ec0..7d0cd046e 100644 --- a/crates/ml/src/dqn/dqn.rs +++ b/crates/ml/src/dqn/dqn.rs @@ -11,6 +11,7 @@ use std::collections::VecDeque; use std::sync::{Arc, Mutex}; +use crate::dqn::mixed_precision::training_dtype; use crate::dqn::target_update::{convergence_half_life, hard_update, polyak_update}; // WAVE 16 (Agent 36) use crate::dqn::xavier_init::linear_xavier; // Xavier initialization with VarMap registration use crate::Adam; @@ -26,6 +27,7 @@ use serde::{Deserialize, Serialize}; use tracing::debug; use super::{Experience, FactoredAction}; +use crate::dqn::mixed_precision::training_dtype; use crate::MLError; /// Configuration for the `DQN` @@ -666,7 +668,7 @@ impl Sequential { _noisy_sigma_init: f64, // Reserved for future custom sigma initialization ) -> Result { let vars = VarMap::new(); - let var_builder = VarBuilder::from_varmap(&vars, DType::F32, &device); + let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); let mut layers = Vec::new(); let mut noisy_layers = Vec::new(); @@ -1540,12 +1542,20 @@ impl DQN { let states_tensor = Tensor::from_vec(states, (batch_size, self.config.state_dim), device) .map_err(|e| { MLError::TrainingError(format!("Failed to create states tensor: {}", e)) + })? + .to_dtype(training_dtype(device)) + .map_err(|e| { + MLError::TrainingError(format!("Failed to cast states tensor to training dtype: {}", e)) })?; let next_states_tensor = Tensor::from_vec(next_states, (batch_size, self.config.state_dim), device).map_err( |e| MLError::TrainingError(format!("Failed to create next states tensor: {}", e)), - )?; + )? + .to_dtype(training_dtype(device)) + .map_err(|e| { + MLError::TrainingError(format!("Failed to cast next states tensor to training dtype: {}", e)) + })?; // Wave 11.7: Validate action indices before creating tensor (CUDA gather bounds check) let num_actions = self.config.num_actions; diff --git a/crates/ml/src/trainers/dqn/trainer.rs b/crates/ml/src/trainers/dqn/trainer.rs index 2f2676cbf..b88c05029 100644 --- a/crates/ml/src/trainers/dqn/trainer.rs +++ b/crates/ml/src/trainers/dqn/trainer.rs @@ -30,6 +30,7 @@ use crate::dqn::portfolio_tracker::PortfolioTracker; use crate::dqn::regime_conditional::RegimeConditionalDQN; use crate::dqn::reward::{RewardConfig, RewardFunction}; use crate::dqn::target_update::convergence_half_life; +use crate::dqn::mixed_precision::training_dtype; use crate::dqn::{Experience, TradingState}; use crate::evaluation::metrics::calculate_var_cvar; use crate::trainers::TargetUpdateMode; @@ -1016,7 +1017,9 @@ impl DQNTrainer { states, (sample_size, state_dim), agent.device() - ).map_err(|e| crate::MLError::ModelError(format!("Failed to create batch tensor: {}", e)))?; + ).map_err(|e| crate::MLError::ModelError(format!("Failed to create batch tensor: {}", e)))? + .to_dtype(training_dtype(agent.device())) + .map_err(|e| crate::MLError::ModelError(format!("Failed to cast batch tensor to training dtype: {}", e)))?; // Forward pass to get Q-values [batch_size, num_actions] let q_values = agent.forward(&batch_tensor)?; @@ -1128,7 +1131,9 @@ impl DQNTrainer { let state_dim = state_vecs[0].len(); let batched_states: Vec = state_vecs.iter().flat_map(|v| v.iter().copied()).collect(); let batch_tensor = Tensor::from_vec(batched_states, (sample_size, state_dim), &self.device) - .map_err(|e| anyhow::anyhow!("Failed to create batched validation tensor: {}", e))?; + .map_err(|e| anyhow::anyhow!("Failed to create batched validation tensor: {}", e))? + .to_dtype(training_dtype(&self.device)) + .map_err(|e| anyhow::anyhow!("Failed to cast validation tensor to training dtype: {}", e))?; let batch_q_values = agent.forward(&batch_tensor) .map_err(|e| anyhow::anyhow!("Batched validation forward pass failed: {}", e))?; @@ -2135,14 +2140,18 @@ impl DQNTrainer { state_vec.clone(), (1, state_vec.len()), &self.device - ).map_err(|e| anyhow::anyhow!("Failed to create state tensor: {}", e))?; + ).map_err(|e| anyhow::anyhow!("Failed to create state tensor: {}", e))? + .to_dtype(training_dtype(&self.device)) + .map_err(|e| anyhow::anyhow!("Failed to cast state tensor to training dtype: {}", e))?; let next_state_vec = next_state.to_vector(); let next_state_tensor = Tensor::from_vec( next_state_vec, (1, state_vec.len()), // Use same length as state_vec &self.device - ).map_err(|e| anyhow::anyhow!("Failed to create next_state tensor: {}", e))?; + ).map_err(|e| anyhow::anyhow!("Failed to create next_state tensor: {}", e))? + .to_dtype(training_dtype(&self.device)) + .map_err(|e| anyhow::anyhow!("Failed to cast next_state tensor to training dtype: {}", e))?; // Calculate intrinsic curiosity reward (prediction error) let intrinsic_reward = curiosity.calculate_curiosity_reward( @@ -3121,7 +3130,9 @@ impl DQNTrainer { // Create batched tensor directly from flat buffer let batch_tensor = Tensor::from_vec(flat_states, (batch_size, state_dim), &self.device) - .map_err(|e| anyhow::anyhow!("Failed to create batched state tensor: {}", e))?; + .map_err(|e| anyhow::anyhow!("Failed to create batched state tensor: {}", e))? + .to_dtype(training_dtype(&self.device)) + .map_err(|e| anyhow::anyhow!("Failed to cast batched state tensor to training dtype: {}", e))?; // WAVE 16S: Get volatility-adjusted epsilon for exploration let base_epsilon = agent.get_epsilon() as f64; @@ -3645,7 +3656,9 @@ impl DQNTrainer { // Create batched tensor [batch_size, state_dim] let batch_tensor = Tensor::from_vec(batched_states, (sample_size, state_dim), agent.device()) - .map_err(|e| anyhow::anyhow!("Failed to create batched state tensor: {}", e))?; + .map_err(|e| anyhow::anyhow!("Failed to create batched state tensor: {}", e))? + .to_dtype(training_dtype(agent.device())) + .map_err(|e| anyhow::anyhow!("Failed to cast batched state tensor to training dtype: {}", e))?; // WAVE 23 P0 Fix: Check for Q-value divergence (early stopping) // This calls log_q_values() which returns Err if divergence detected for consecutive checks