diff --git a/crates/ml/src/ensemble/adapters/diffusion.rs b/crates/ml/src/ensemble/adapters/diffusion.rs index 2129b22b1..91d2ec9cf 100644 --- a/crates/ml/src/ensemble/adapters/diffusion.rs +++ b/crates/ml/src/ensemble/adapters/diffusion.rs @@ -6,11 +6,12 @@ use std::sync::Mutex; -use candle_core::{DType, Device, Tensor}; +use candle_core::{Device, Tensor}; use candle_nn::{VarBuilder, VarMap}; use crate::diffusion::config::DiffusionConfig; use crate::diffusion::denoiser::Denoiser; +use crate::dqn::mixed_precision::training_dtype; use crate::ensemble::inference_adapter::{ EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, RawPrediction, }; @@ -46,7 +47,7 @@ impl DiffusionInferenceAdapter { let device = DeviceConfig::Auto.resolve().unwrap_or(Device::Cpu); let data_dim = config.seq_len * config.feature_dim; let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let denoiser = Denoiser::new( data_dim, config.hidden_dim, @@ -69,7 +70,7 @@ impl DiffusionInferenceAdapter { let device = DeviceConfig::Auto.resolve().unwrap_or(Device::Cpu); let data_dim = config.seq_len * config.feature_dim; let mut varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let denoiser = Denoiser::new( data_dim, config.hidden_dim, diff --git a/crates/ml/src/ensemble/adapters/kan.rs b/crates/ml/src/ensemble/adapters/kan.rs index 7a57a0f7c..7868fa79e 100644 --- a/crates/ml/src/ensemble/adapters/kan.rs +++ b/crates/ml/src/ensemble/adapters/kan.rs @@ -5,9 +5,10 @@ use std::sync::Mutex; -use candle_core::{DType, Device, Tensor}; +use candle_core::{Device, Tensor}; use candle_nn::{VarBuilder, VarMap}; +use crate::dqn::mixed_precision::training_dtype; use crate::ensemble::inference_adapter::{ EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, RawPrediction, }; @@ -41,7 +42,7 @@ impl KanInferenceAdapter { let device = DeviceConfig::Auto.resolve().unwrap_or(Device::Cpu); let input_dim = config.layer_widths.first().copied().unwrap_or(51); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let network = KANNetwork::new(&config, vb)?; Ok(Self { @@ -57,7 +58,7 @@ impl KanInferenceAdapter { let device = DeviceConfig::Auto.resolve().unwrap_or(Device::Cpu); let input_dim = config.layer_widths.first().copied().unwrap_or(51); let mut varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let network = KANNetwork::new(&config, vb)?; varmap diff --git a/crates/ml/src/ensemble/adapters/liquid.rs b/crates/ml/src/ensemble/adapters/liquid.rs index fbc04f16c..1119fca8d 100644 --- a/crates/ml/src/ensemble/adapters/liquid.rs +++ b/crates/ml/src/ensemble/adapters/liquid.rs @@ -5,9 +5,10 @@ use std::sync::Mutex; -use candle_core::{DType, Device, Tensor}; +use candle_core::{Device, Tensor}; use candle_nn::{VarBuilder, VarMap}; +use crate::dqn::mixed_precision::training_dtype; use crate::ensemble::inference_adapter::{ EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, }; @@ -43,7 +44,7 @@ impl LiquidInferenceAdapter { let device = config.device.resolve().unwrap_or(Device::Cpu); let input_size = config.input_size; let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let network = CandleCfCNetwork::new(&config, &vb)?; Ok(Self { @@ -59,7 +60,7 @@ impl LiquidInferenceAdapter { let device = config.device.resolve().unwrap_or(Device::Cpu); let input_size = config.input_size; let mut varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let network = CandleCfCNetwork::new(&config, &vb)?; varmap diff --git a/crates/ml/src/ensemble/adapters/tggn.rs b/crates/ml/src/ensemble/adapters/tggn.rs index 2168f63c1..415bab74b 100644 --- a/crates/ml/src/ensemble/adapters/tggn.rs +++ b/crates/ml/src/ensemble/adapters/tggn.rs @@ -5,9 +5,10 @@ use std::sync::Mutex; -use candle_core::{DType, Device, Module, Tensor}; +use candle_core::{Device, Module, Tensor}; use candle_nn::{linear, Linear, VarBuilder, VarMap}; +use crate::dqn::mixed_precision::training_dtype; use crate::ensemble::inference_adapter::{ EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, RawPrediction, }; @@ -74,7 +75,7 @@ impl TggnInferenceAdapter { pub fn new(input_dim: usize, hidden_dim: usize) -> MLResult { let device = DeviceConfig::Auto.resolve().unwrap_or(Device::Cpu); let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let projection = TggnProjection::new(input_dim, hidden_dim, vb)?; Ok(Self { @@ -93,7 +94,7 @@ impl TggnInferenceAdapter { ) -> MLResult { let device = DeviceConfig::Auto.resolve().unwrap_or(Device::Cpu); let mut varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let projection = TggnProjection::new(input_dim, hidden_dim, vb)?; varmap diff --git a/crates/ml/src/ensemble/adapters/tlob.rs b/crates/ml/src/ensemble/adapters/tlob.rs index 066051942..b5e1d5254 100644 --- a/crates/ml/src/ensemble/adapters/tlob.rs +++ b/crates/ml/src/ensemble/adapters/tlob.rs @@ -9,9 +9,10 @@ use std::collections::VecDeque; use std::sync::Mutex; -use candle_core::{DType, Device, Module, Tensor}; +use candle_core::{Device, Module, Tensor}; use candle_nn::{linear, Linear, VarBuilder, VarMap}; +use crate::dqn::mixed_precision::training_dtype; use crate::ensemble::inference_adapter::{ EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, RawPrediction, }; @@ -97,7 +98,7 @@ impl TlobInferenceAdapter { let device = DeviceConfig::Auto.resolve().unwrap_or(Device::Cpu); let flat_dim = sequence_length * feature_dim; let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let projection = TlobProjection::new(flat_dim, hidden_dim, vb)?; Ok(Self { @@ -120,7 +121,7 @@ impl TlobInferenceAdapter { let device = DeviceConfig::Auto.resolve().unwrap_or(Device::Cpu); let flat_dim = sequence_length * feature_dim; let mut varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let projection = TlobProjection::new(flat_dim, hidden_dim, vb)?; varmap diff --git a/crates/ml/src/ensemble/adapters/xlstm.rs b/crates/ml/src/ensemble/adapters/xlstm.rs index 6b9a6202d..1539eb726 100644 --- a/crates/ml/src/ensemble/adapters/xlstm.rs +++ b/crates/ml/src/ensemble/adapters/xlstm.rs @@ -8,9 +8,10 @@ use std::collections::VecDeque; use std::sync::Mutex; -use candle_core::{DType, Device, Tensor}; +use candle_core::{Device, Tensor}; use candle_nn::{VarBuilder, VarMap}; +use crate::dqn::mixed_precision::training_dtype; use crate::ensemble::inference_adapter::{ EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, RawPrediction, }; @@ -54,7 +55,7 @@ impl XlstmInferenceAdapter { let device = DeviceConfig::Auto.resolve().unwrap_or(Device::Cpu); let input_dim = config.input_dim; let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let network = XLSTMNetwork::new(&config, vb)?; Ok(Self { @@ -76,7 +77,7 @@ impl XlstmInferenceAdapter { let device = DeviceConfig::Auto.resolve().unwrap_or(Device::Cpu); let input_dim = config.input_dim; let mut varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let network = XLSTMNetwork::new(&config, vb)?; varmap diff --git a/crates/ml/src/explainability/integrated_gradients.rs b/crates/ml/src/explainability/integrated_gradients.rs index 9cdf42145..bc359ec60 100644 --- a/crates/ml/src/explainability/integrated_gradients.rs +++ b/crates/ml/src/explainability/integrated_gradients.rs @@ -168,7 +168,7 @@ impl IntegratedGradients { #[cfg(test)] mod tests { use super::*; - use candle_core::{DType, Device}; + use candle_core::Device; use candle_nn::{linear, VarBuilder, VarMap}; /// A simple 2-layer test network: linear(4,8) -> relu -> linear(8,1) @@ -197,7 +197,7 @@ mod tests { fn test_integrated_gradients_basic() { let device = Device::Cpu; let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vs = VarBuilder::from_varmap(&varmap, crate::dqn::mixed_precision::training_dtype(&device), &device); let model = TwoLayerNet::new(vs); // Linear layers require 2D input: [batch, features] @@ -244,7 +244,7 @@ mod tests { fn test_ig_completeness_axiom() { let device = Device::Cpu; let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vs = VarBuilder::from_varmap(&varmap, crate::dqn::mixed_precision::training_dtype(&device), &device); let model = LinearNet::new(vs); let input = Tensor::new(&[[1.0_f32, -0.5, 0.3, 2.0]], &device).unwrap(); @@ -298,7 +298,7 @@ mod tests { fn test_ig_dimension_mismatch() { let device = Device::Cpu; let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vs = VarBuilder::from_varmap(&varmap, crate::dqn::mixed_precision::training_dtype(&device), &device); let model = TwoLayerNet::new(vs); let input = Tensor::new(&[[1.0_f32, -0.5, 0.3, 2.0]], &device).unwrap(); diff --git a/crates/ml/src/features/multi_timeframe.rs b/crates/ml/src/features/multi_timeframe.rs index 08d053418..7ca378384 100644 --- a/crates/ml/src/features/multi_timeframe.rs +++ b/crates/ml/src/features/multi_timeframe.rs @@ -18,6 +18,7 @@ use candle_core::{DType, Device, Tensor}; use candle_nn::{linear, Linear, Module, VarBuilder, VarMap}; use super::bar_resampler::BarResampler; +use crate::dqn::mixed_precision::training_dtype; use crate::types::OHLCVBar; use crate::MLError; @@ -266,7 +267,7 @@ impl MultiTimeframeEncoder { device: &Device, ) -> Result<(Self, VarMap), MLError> { let vars = VarMap::new(); - let vb = VarBuilder::from_varmap(&vars, DType::F32, device); + let vb = VarBuilder::from_varmap(&vars, training_dtype(device), device); let encoder = Self::new(config, vb)?; Ok((encoder, vars)) } @@ -557,7 +558,7 @@ mod tests { #[test] fn test_lstm_encoder_single_step() { let vars = VarMap::new(); - let vb = VarBuilder::from_varmap(&vars, DType::F32, &Device::Cpu); + let vb = VarBuilder::from_varmap(&vars, training_dtype(&Device::Cpu), &Device::Cpu); let lstm = LstmEncoder::new(6, 32, vb.pp("test_lstm")).expect("lstm creation"); // Single timestep: (1, 6) @@ -573,7 +574,7 @@ mod tests { #[test] fn test_lstm_encoder_multi_step() { let vars = VarMap::new(); - let vb = VarBuilder::from_varmap(&vars, DType::F32, &Device::Cpu); + let vb = VarBuilder::from_varmap(&vars, training_dtype(&Device::Cpu), &Device::Cpu); let lstm = LstmEncoder::new(6, 64, vb.pp("test_lstm")).expect("lstm creation"); // 10 timesteps: (10, 6) diff --git a/crates/ml/src/portfolio_transformer.rs b/crates/ml/src/portfolio_transformer.rs index afc8e77ba..bf6de2a42 100644 --- a/crates/ml/src/portfolio_transformer.rs +++ b/crates/ml/src/portfolio_transformer.rs @@ -4,13 +4,14 @@ //! in high-frequency trading. Unlike traditional time-series transformers, this model //! operates directly on portfolio state vectors for optimal weight prediction. -use candle_core::{DType, Device, IndexOp, Module, ModuleT, Result as CandleResult, Tensor}; +use candle_core::{Device, IndexOp, Module, ModuleT, Result as CandleResult, Tensor}; use candle_nn::{Linear, VarBuilder, VarMap}; use chrono::{DateTime, Utc}; use serde::{Deserialize, Serialize}; use tracing::{debug, instrument, warn}; use super::*; +use crate::dqn::mixed_precision::training_dtype; /// Portfolio state representation for transformer input #[derive(Debug, Clone, Serialize, Deserialize)] @@ -187,7 +188,7 @@ impl PortfolioTransformer { /// Create new Portfolio Transformer pub fn new(config: PortfolioTransformerConfig, device: Device) -> MLResult { let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vb = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); // Input projection let input_projection = candle_nn::linear( diff --git a/crates/ml/src/ppo/continuous_policy.rs b/crates/ml/src/ppo/continuous_policy.rs index 909b64eea..69a523bcb 100644 --- a/crates/ml/src/ppo/continuous_policy.rs +++ b/crates/ml/src/ppo/continuous_policy.rs @@ -12,7 +12,7 @@ use std::f32::consts::PI; -use candle_core::{DType, Device, Tensor}; +use candle_core::{Device, Tensor}; use candle_nn::{linear, Linear, Module, VarBuilder, VarMap}; use rand::thread_rng; // Note: rand_distr::Distribution could be used for direct sampling if added to dependencies @@ -21,6 +21,7 @@ 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; @@ -80,7 +81,7 @@ impl ContinuousPolicyNetwork { /// Create new continuous policy network pub fn new(config: ContinuousPolicyConfig, device: Device) -> 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 feature_layers = Vec::new(); let mut current_dim = config.state_dim; diff --git a/crates/ml/src/ppo/flow_policy/mod.rs b/crates/ml/src/ppo/flow_policy/mod.rs index 06b4079ff..ef88e0534 100644 --- a/crates/ml/src/ppo/flow_policy/mod.rs +++ b/crates/ml/src/ppo/flow_policy/mod.rs @@ -9,6 +9,7 @@ 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; @@ -121,7 +122,7 @@ impl FlowPolicy { /// A new FlowPolicy instance or an error if initialization fails. pub fn new(config: FlowPolicyConfig, device: &Device) -> Result { let vars = VarMap::new(); - let vb = VarBuilder::from_varmap(&vars, DType::F32, device); + let vb = VarBuilder::from_varmap(&vars, training_dtype(device), device); // Context encoder: state_dim → context_dim let context_enc = linear_xavier( diff --git a/crates/ml/src/ppo/lstm_networks.rs b/crates/ml/src/ppo/lstm_networks.rs index 316abd6cd..7b69fbb73 100644 --- a/crates/ml/src/ppo/lstm_networks.rs +++ b/crates/ml/src/ppo/lstm_networks.rs @@ -5,10 +5,11 @@ //! these networks maintain hidden states across timesteps, enabling the agent to //! "remember" past observations when making decisions. -use candle_core::{DType, Device, Tensor}; +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; /// LSTM-augmented policy network for temporal action selection @@ -49,7 +50,7 @@ impl LSTMPolicyNetwork { device: Device, ) -> 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); // Input projection layer (state_dim → hidden_dim) let input_layer = linear(input_dim, hidden_dim, var_builder.pp("input")) @@ -264,7 +265,7 @@ impl LSTMValueNetwork { device: Device, ) -> 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); // Input projection layer (state_dim → hidden_dim) let input_layer = linear(input_dim, hidden_dim, var_builder.pp("input")) diff --git a/crates/ml/src/ppo/ppo.rs b/crates/ml/src/ppo/ppo.rs index d843c15b5..7c34d9ae2 100644 --- a/crates/ml/src/ppo/ppo.rs +++ b/crates/ml/src/ppo/ppo.rs @@ -28,6 +28,7 @@ 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::dqn::reward::RewardNormalizer; @@ -298,7 +299,7 @@ impl PolicyNetwork { device: Device, ) -> 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 current_dim = input_dim; @@ -546,7 +547,7 @@ impl ValueNetwork { /// Create new value network pub fn new(input_dim: usize, hidden_dims: &[usize], device: Device) -> 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 current_dim = input_dim; @@ -1792,7 +1793,7 @@ impl PPO { // 4. Candle's deserializer validates format before tensor creation // 5. Any format violations cause Err return, not UB let actor_vb = unsafe { - VarBuilder::from_mmaped_safetensors(&[actor_path], DType::F32, &device).map_err( + VarBuilder::from_mmaped_safetensors(&[actor_path], training_dtype(&device), &device).map_err( |e| { MLError::ModelError(format!( "Failed to load actor checkpoint from {}: {}", @@ -1849,7 +1850,7 @@ impl PPO { // 4. Candle's deserializer validates format before tensor creation // 5. Any format violations cause Err return, not UB let critic_vb = unsafe { - VarBuilder::from_mmaped_safetensors(&[critic_path], DType::F32, &device).map_err( + VarBuilder::from_mmaped_safetensors(&[critic_path], training_dtype(&device), &device).map_err( |e| { MLError::ModelError(format!( "Failed to load critic checkpoint from {}: {}", diff --git a/crates/ml/src/trainers/online_learning.rs b/crates/ml/src/trainers/online_learning.rs index cb47b1fbc..174f6d76d 100644 --- a/crates/ml/src/trainers/online_learning.rs +++ b/crates/ml/src/trainers/online_learning.rs @@ -555,7 +555,6 @@ impl OnlineLearner { #[cfg(test)] mod tests { use super::*; - use candle_core::DType; // ----------------------------------------------------------------------- // Helpers @@ -578,7 +577,7 @@ mod tests { /// Build a tiny VarMap with a single 2x2 parameter for testing. fn tiny_var_map() -> Result { let var_map = VarMap::new(); - let vb = candle_nn::VarBuilder::from_varmap(&var_map, DType::F32, &Device::Cpu); + let vb = candle_nn::VarBuilder::from_varmap(&var_map, crate::dqn::mixed_precision::training_dtype(&Device::Cpu), &Device::Cpu); let _linear = candle_nn::linear(2, 2, vb.pp("layer")) .map_err(|e| MLError::ModelError(format!("tiny_var_map linear: {e}")))?; Ok(var_map) diff --git a/crates/ml/src/trainers/ppo.rs b/crates/ml/src/trainers/ppo.rs index b260af645..43e49888e 100644 --- a/crates/ml/src/trainers/ppo.rs +++ b/crates/ml/src/trainers/ppo.rs @@ -25,6 +25,7 @@ use crate::ppo::gae::GAEConfig; use crate::ppo::ppo::{PPOConfig, PPO}; use crate::ppo::trajectories::{Trajectory, TrajectoryBatch, TrajectoryStep}; use crate::cuda_pipeline::PpoGpuData; +use crate::dqn::mixed_precision::training_dtype; use crate::MLError; /// PPO training hyperparameters (matches gRPC PpoParams) @@ -832,6 +833,7 @@ impl PpoTrainer { .map_err(|e| MLError::ModelError(format!("PPO GPU state access failed: {e}")))? } else { Tensor::from_vec(state.clone(), (1, state.len()), &self.device)? + .to_dtype(training_dtype(&self.device))? }; // Actor forward pass — use pre-uploaded tensor @@ -892,7 +894,8 @@ impl PpoTrainer { // Phase 2: Batched critic forward — SINGLE GPU→CPU sync for all value estimates if step_count > 0 && state_dim > 0 { let all_states_tensor = - Tensor::from_vec(all_state_floats, (step_count, state_dim), &self.device)?; + Tensor::from_vec(all_state_floats, (step_count, state_dim), &self.device)? + .to_dtype(training_dtype(&self.device))?; let all_values_vec = model .critic .forward(&all_states_tensor)? @@ -958,6 +961,7 @@ impl PpoTrainer { .map_err(|e| MLError::ModelError(format!("PPO GPU state access failed: {e}")))? } else { Tensor::from_vec(state.clone(), (1, state.len()), &self.device)? + .to_dtype(training_dtype(&self.device))? }; // Get action from policy — use pre-uploaded tensor @@ -1011,7 +1015,8 @@ impl PpoTrainer { // Phase 2: Batched critic forward — SINGLE GPU→CPU sync for all value estimates if step_count > 0 && state_dim > 0 { let all_states_tensor = - Tensor::from_vec(all_state_floats, (step_count, state_dim), &self.device)?; + Tensor::from_vec(all_state_floats, (step_count, state_dim), &self.device)? + .to_dtype(training_dtype(&self.device))?; let all_values_vec = model .critic .forward(&all_states_tensor)?