diff --git a/crates/ml/src/dqn/agent.rs b/crates/ml/src/dqn/agent.rs index 10f76db02..e3b81b2a2 100644 --- a/crates/ml/src/dqn/agent.rs +++ b/crates/ml/src/dqn/agent.rs @@ -8,6 +8,7 @@ use std::collections::HashMap; use crate::Adam; use candle_core::Tensor; use candle_nn::{ops::leaky_relu, Module, VarBuilder}; +use crate::dqn::mixed_precision::training_dtype; use candle_optimisers::adam::ParamsAdam; // Use our Adam wrapper from lib.rs use serde::{Deserialize, Serialize}; use tracing::debug; @@ -363,12 +364,12 @@ impl DQNAgent { // Forward pass through main network with gradient tracking let var_builder = - VarBuilder::from_varmap(self.q_network.vars(), candle_core::DType::F32, device); + VarBuilder::from_varmap(self.q_network.vars(), training_dtype(device), device); let current_q_values = self.forward_with_gradients(&state_tensor, &var_builder)?; // Forward pass through target network WITHOUT gradients let target_var_builder = - VarBuilder::from_varmap(self.target_network.vars(), candle_core::DType::F32, device); + VarBuilder::from_varmap(self.target_network.vars(), training_dtype(device), device); let next_q_values = self.forward_without_gradients(&next_state_tensor, &target_var_builder)?; @@ -594,7 +595,7 @@ impl DQNAgent { self.q_network.vars() }; let var_builder = - VarBuilder::from_varmap(vars, candle_core::DType::F32, self.q_network.device()); + VarBuilder::from_varmap(vars, training_dtype(self.q_network.device()), self.q_network.device()); // Reconstruct network layers let mut layers = Vec::new(); diff --git a/crates/ml/src/dqn/attention.rs b/crates/ml/src/dqn/attention.rs index 49e03965c..2be4e7142 100644 --- a/crates/ml/src/dqn/attention.rs +++ b/crates/ml/src/dqn/attention.rs @@ -423,8 +423,8 @@ impl MultiHeadAttention { #[cfg(test)] mod tests { use super::*; - use candle_core::DType; use candle_nn::VarMap; + use crate::dqn::mixed_precision::training_dtype; #[test] fn test_config_validation() { @@ -462,7 +462,7 @@ mod tests { let device = Device::Cpu; let config = MultiHeadAttentionConfig::new(64, 4)?; let vars = VarMap::new(); - let vb = VarBuilder::from_varmap(&vars, DType::F32, &device); + let vb = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); let attention = MultiHeadAttention::new(config, &vb, &device)?; assert_eq!(attention.config().embed_dim, 64); @@ -476,7 +476,7 @@ mod tests { let device = Device::Cpu; let config = MultiHeadAttentionConfig::new(64, 4)?; let vars = VarMap::new(); - let vb = VarBuilder::from_varmap(&vars, DType::F32, &device); + let vb = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); let attention = MultiHeadAttention::new(config, &vb, &device)?; @@ -506,7 +506,7 @@ mod tests { let device = Device::Cpu; let config = MultiHeadAttentionConfig::new(64, 4)?; let vars = VarMap::new(); - let vb = VarBuilder::from_varmap(&vars, DType::F32, &device); + let vb = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); let attention = MultiHeadAttention::new(config, &vb, &device)?; @@ -546,7 +546,7 @@ mod tests { let device = Device::Cpu; let config = MultiHeadAttentionConfig::new(64, 4)?; let vars = VarMap::new(); - let vb = VarBuilder::from_varmap(&vars, DType::F32, &device); + let vb = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); let attention = MultiHeadAttention::new(config, &vb, &device)?; @@ -580,7 +580,7 @@ mod tests { config.use_layer_norm = false; // Disable to test residual alone let vars = VarMap::new(); - let vb = VarBuilder::from_varmap(&vars, DType::F32, &device); + let vb = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); let attention = MultiHeadAttention::new(config, &vb, &device)?; @@ -611,7 +611,7 @@ mod tests { let embed_dim = 64; let config = MultiHeadAttentionConfig::new(embed_dim, num_heads)?; let vars = VarMap::new(); - let vb = VarBuilder::from_varmap(&vars, DType::F32, &device); + let vb = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); let attention = MultiHeadAttention::new(config, &vb, &device)?; diff --git a/crates/ml/src/dqn/curiosity.rs b/crates/ml/src/dqn/curiosity.rs index 27d3d78e7..3e2c07c19 100644 --- a/crates/ml/src/dqn/curiosity.rs +++ b/crates/ml/src/dqn/curiosity.rs @@ -7,6 +7,7 @@ use candle_core::{DType, Device, Tensor}; use candle_nn::{ops::leaky_relu, AdamW, Linear, Module, Optimizer, ParamsAdamW, VarBuilder, VarMap}; use super::action_space::{FactoredAction, ExposureLevel}; +use super::mixed_precision::training_dtype; use crate::MLError; use crate::dqn::xavier_init::linear_xavier; @@ -35,7 +36,7 @@ impl ForwardDynamicsModel { /// - Output: 32 (predicted next state embedding) fn new(device: Device, _learning_rate: f64) -> 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: 32 state + 3 action one-hot = 35 // Hidden: 64 diff --git a/crates/ml/src/dqn/distributional_dueling.rs b/crates/ml/src/dqn/distributional_dueling.rs index 8e4ab7be5..95f88eab6 100644 --- a/crates/ml/src/dqn/distributional_dueling.rs +++ b/crates/ml/src/dqn/distributional_dueling.rs @@ -43,6 +43,7 @@ use candle_core::{DType, Device, Tensor}; use candle_nn::{Linear, Module, VarBuilder, VarMap}; use serde::{Deserialize, Serialize}; +use crate::dqn::mixed_precision::training_dtype; use crate::dqn::xavier_init::linear_xavier; use crate::MLError; @@ -152,7 +153,7 @@ impl DistributionalDuelingQNetwork { /// New DistributionalDuelingQNetwork instance with Xavier-initialized weights pub fn new(config: DistributionalDuelingConfig, 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); // Build shared feature layers let mut shared_layers = Vec::new(); diff --git a/crates/ml/src/dqn/dqn.rs b/crates/ml/src/dqn/dqn.rs index 7d0cd046e..346c84d13 100644 --- a/crates/ml/src/dqn/dqn.rs +++ b/crates/ml/src/dqn/dqn.rs @@ -27,7 +27,6 @@ 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` diff --git a/crates/ml/src/dqn/dueling.rs b/crates/ml/src/dqn/dueling.rs index b0b6a5f24..fe08c228e 100644 --- a/crates/ml/src/dqn/dueling.rs +++ b/crates/ml/src/dqn/dueling.rs @@ -34,6 +34,7 @@ use candle_core::{DType, Device, Tensor}; use candle_nn::{Linear, Module, VarBuilder, VarMap}; use serde::{Deserialize, Serialize}; +use crate::dqn::mixed_precision::training_dtype; use crate::dqn::xavier_init::linear_xavier; use crate::MLError; @@ -136,7 +137,7 @@ impl DuelingQNetwork { /// New DuelingQNetwork instance with Xavier-initialized weights pub fn new(config: DuelingConfig, 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); // Build shared feature layers let mut shared_layers = Vec::new(); diff --git a/crates/ml/src/dqn/factored_q_network.rs b/crates/ml/src/dqn/factored_q_network.rs index 89f66e5de..4f74c930c 100644 --- a/crates/ml/src/dqn/factored_q_network.rs +++ b/crates/ml/src/dqn/factored_q_network.rs @@ -13,6 +13,8 @@ use candle_core::{Device, Tensor}; use candle_nn::{ops::leaky_relu, Linear, Module, VarBuilder, VarMap}; use rand::Rng; + +use crate::dqn::mixed_precision::training_dtype; use serde::{Deserialize, Serialize}; use super::action_space::{ExposureLevel, FactoredAction, OrderType, Urgency}; @@ -69,7 +71,7 @@ impl FactoredQNetwork { /// Create a new factored Q-network with custom configuration pub fn with_config(config: FactoredQNetworkConfig, device: &Device) -> Result { let varmap = VarMap::new(); - let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::F32, device); + let vb = VarBuilder::from_varmap(&varmap, training_dtype(device), device); // Initialize shared encoder with Xavier uniform let shared_encoder = linear_xavier( diff --git a/crates/ml/src/dqn/network.rs b/crates/ml/src/dqn/network.rs index dcad02baf..5082d62f9 100644 --- a/crates/ml/src/dqn/network.rs +++ b/crates/ml/src/dqn/network.rs @@ -3,6 +3,7 @@ use std::sync::atomic::{AtomicU32, AtomicU64, Ordering}; use candle_core::{DType, Device, Result as CandleResult, Tensor}; +use crate::dqn::mixed_precision::training_dtype; use candle_nn::Module; use candle_nn::{ops::leaky_relu, Dropout, Linear, VarBuilder, VarMap}; use rand::prelude::*; // Replace common::rng with standard rand @@ -253,12 +254,12 @@ impl QNetwork { let target_vars = VarMap::new(); // Initialize network weights - let var_builder = VarBuilder::from_varmap(&vars, DType::F32, &device); + let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); let _layers = NetworkLayers::new(&var_builder, &config, &device) .map_err(|e| MLError::ModelError(format!("Failed to create network layers: {}", e)))?; // Initialize target network with same architecture - let target_var_builder = VarBuilder::from_varmap(&target_vars, DType::F32, &device); + let target_var_builder = VarBuilder::from_varmap(&target_vars, training_dtype(&device), &device); let _target_layers = NetworkLayers::new(&target_var_builder, &config, &device).map_err(|e| { MLError::ModelError(format!("Failed to create target network layers: {}", e)) @@ -298,7 +299,7 @@ impl QNetwork { // Get current dropout rate (adaptive or static) let dropout_rate = self.get_dropout_rate(); - let var_builder = VarBuilder::from_varmap(&self.vars, DType::F32, &self.device); + let var_builder = VarBuilder::from_varmap(&self.vars, training_dtype(&self.device), &self.device); let layers = NetworkLayers::new_with_dropout_rate( &var_builder, &self.config, @@ -360,7 +361,7 @@ impl QNetwork { flat_states.extend_from_slice(state); } - let var_builder = VarBuilder::from_varmap(&self.vars, DType::F32, &self.device); + let var_builder = VarBuilder::from_varmap(&self.vars, training_dtype(&self.device), &self.device); let layers = NetworkLayers::new(&var_builder, &self.config, &self.device) .map_err(|e| MLError::ModelError(format!("Failed to create layers: {}", e)))?; diff --git a/crates/ml/src/dqn/noisy_layers.rs b/crates/ml/src/dqn/noisy_layers.rs index 8ac8d0026..2f7be04fd 100644 --- a/crates/ml/src/dqn/noisy_layers.rs +++ b/crates/ml/src/dqn/noisy_layers.rs @@ -304,14 +304,14 @@ impl Default for NoisyNetworkConfig { #[cfg(test)] mod tests { use super::*; - use candle_core::DType; use candle_nn::{VarBuilder, VarMap}; + use crate::dqn::mixed_precision::training_dtype; #[test] fn test_noisy_linear_creation() -> Result<(), MLError> { let device = 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 _layer = NoisyLinear::new(64, 32, vb)?; Ok(()) @@ -321,7 +321,7 @@ mod tests { fn test_noisy_linear_forward() -> Result<(), MLError> { let device = 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 mut layer = NoisyLinear::new(64, 32, vb)?; layer.reset_noise()?; // Resample noise before forward @@ -343,7 +343,7 @@ mod tests { fn test_noise_reset() -> Result<(), MLError> { let device = 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 mut layer = NoisyLinear::new(64, 32, vb)?; let input = Tensor::randn(0.0_f32, 1.0_f32, (4, 64), &device) @@ -388,7 +388,7 @@ mod tests { fn test_disable_noise() -> Result<(), MLError> { let device = 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 mut layer = NoisyLinear::new(64, 32, vb)?; let input = Tensor::randn(0.0_f32, 1.0_f32, (4, 64), &device) @@ -428,7 +428,7 @@ mod tests { fn test_factorized_noise_dimensions() -> Result<(), MLError> { let device = 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 mut layer = NoisyLinear::new(128, 64, vb)?; layer.reset_noise()?; @@ -444,7 +444,7 @@ mod tests { fn test_reset_noise_with_sigma() -> Result<(), MLError> { let device = 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 mut layer = NoisyLinear::new(64, 32, vb)?; let input = Tensor::randn(0.0_f32, 1.0_f32, (4, 64), &device) @@ -484,7 +484,7 @@ mod tests { fn test_sigma_scaling_effect() -> Result<(), MLError> { let device = 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 mut layer = NoisyLinear::new(64, 32, vb)?; diff --git a/crates/ml/src/dqn/quantile_regression.rs b/crates/ml/src/dqn/quantile_regression.rs index 154bb6e47..1d1f86da4 100644 --- a/crates/ml/src/dqn/quantile_regression.rs +++ b/crates/ml/src/dqn/quantile_regression.rs @@ -15,11 +15,12 @@ //! 3. **Flexibility**: No need to specify value ranges (v_min/v_max) //! 4. **Stability**: Quantile Huber loss is more robust than cross-entropy -use candle_core::{Device, Result as CandleResult, Tensor, DType}; +use candle_core::{DType, Device, Result as CandleResult, Tensor}; use candle_nn::{Linear, Module, VarBuilder, VarMap}; use serde::{Deserialize, Serialize}; use std::f32::consts::PI; +use crate::dqn::mixed_precision::training_dtype; use crate::MLError; /// Configuration for Quantile Regression DQN @@ -87,7 +88,7 @@ impl QuantileNetwork { vars: VarMap, device: &Device, ) -> Result { - let vb = VarBuilder::from_varmap(&vars, DType::F32, device); + let vb = VarBuilder::from_varmap(&vars, training_dtype(device), device); // Quantile embedding layer let quantile_embedding = candle_nn::linear( diff --git a/crates/ml/src/dqn/rainbow_agent.rs b/crates/ml/src/dqn/rainbow_agent.rs index 98fc38346..10c4b8236 100644 --- a/crates/ml/src/dqn/rainbow_agent.rs +++ b/crates/ml/src/dqn/rainbow_agent.rs @@ -11,8 +11,10 @@ use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::Arc; -use candle_core::{DType, Device}; +use candle_core::Device; use candle_nn::{VarBuilder, VarMap}; + +use crate::dqn::mixed_precision::training_dtype; use parking_lot::Mutex; use super::*; @@ -38,11 +40,11 @@ impl RainbowAgent { // Create VarMap and VarBuilder for network initialization let varmap = VarMap::new(); - let _vs = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let _vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); // Create VarMap and VarBuilder for network let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let network = RainbowNetwork::new(&vs, config.network_config.clone())?; Ok(Self { config, diff --git a/crates/ml/src/dqn/rainbow_agent_impl.rs b/crates/ml/src/dqn/rainbow_agent_impl.rs index 449c382f5..e44c0749d 100644 --- a/crates/ml/src/dqn/rainbow_agent_impl.rs +++ b/crates/ml/src/dqn/rainbow_agent_impl.rs @@ -8,8 +8,9 @@ use std::collections::VecDeque; use std::sync::{Arc, Mutex, RwLock}; use crate::Adam; -use candle_core::{DType, Device, Tensor}; +use candle_core::{Device, Tensor}; use candle_nn::{VarBuilder, VarMap}; +use crate::dqn::mixed_precision::training_dtype; use candle_optimisers::adam::ParamsAdam; use tracing::{debug, info}; @@ -66,10 +67,10 @@ impl RainbowAgent { let target_varmap = Arc::new(VarMap::new()); // Create networks - let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let online_network = RainbowNetwork::new(&vs, config.network_config.clone())?; - let target_vs = VarBuilder::from_varmap(&target_varmap, DType::F32, &device); + let target_vs = VarBuilder::from_varmap(&target_varmap, training_dtype(&device), &device); let target_network = RainbowNetwork::new(&target_vs, config.network_config.clone())?; // Create optimizer diff --git a/crates/ml/src/dqn/rainbow_network.rs b/crates/ml/src/dqn/rainbow_network.rs index 0658557a7..d3015c8bb 100644 --- a/crates/ml/src/dqn/rainbow_network.rs +++ b/crates/ml/src/dqn/rainbow_network.rs @@ -419,14 +419,15 @@ impl Module for RainbowNetwork { mod tests { use super::*; use anyhow::Result; - use candle_core::{DType, Device}; + use candle_core::Device; use candle_nn::{VarBuilder, VarMap}; + use crate::dqn::mixed_precision::training_dtype; #[test] fn test_rainbow_network_creation() -> Result<(), MLError> { let device = Device::Cpu; let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let config = RainbowNetworkConfig::default(); let _network = RainbowNetwork::new(&vs, config) @@ -447,7 +448,7 @@ mod tests { fn test_rainbow_activation_types() -> Result<(), MLError> { let device = Device::Cpu; let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let mut config = RainbowNetworkConfig::default(); config.activation = ActivationType::ReLU; diff --git a/crates/ml/src/dqn/residual.rs b/crates/ml/src/dqn/residual.rs index ff91c179c..1a5ee9cbd 100644 --- a/crates/ml/src/dqn/residual.rs +++ b/crates/ml/src/dqn/residual.rs @@ -162,6 +162,7 @@ mod tests { use super::*; use candle_core::{DType, Device}; use candle_nn::VarMap; + use crate::dqn::mixed_precision::training_dtype; #[test] fn test_residual_config_default() { @@ -175,7 +176,7 @@ mod tests { fn test_residual_block_creation() -> anyhow::Result<()> { let device = Device::Cpu; 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 config = ResidualConfig { hidden_dim: 64, @@ -193,7 +194,7 @@ mod tests { fn test_residual_block_forward_train() -> anyhow::Result<()> { let device = Device::Cpu; 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 config = ResidualConfig { hidden_dim: 32, @@ -219,7 +220,7 @@ mod tests { fn test_residual_block_forward_eval() -> anyhow::Result<()> { let device = Device::Cpu; 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 config = ResidualConfig { hidden_dim: 32, @@ -246,7 +247,7 @@ mod tests { // Test that skip connection preserves gradient flow let device = Device::Cpu; 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 config = ResidualConfig { hidden_dim: 16, @@ -273,7 +274,7 @@ mod tests { fn test_residual_batch_processing() -> anyhow::Result<()> { let device = Device::Cpu; 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 config = ResidualConfig { hidden_dim: 64, @@ -298,7 +299,7 @@ mod tests { // Test that gradients can flow through skip connection let device = Device::Cpu; 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 config = ResidualConfig { hidden_dim: 8, @@ -327,7 +328,7 @@ mod tests { fn test_residual_different_dimensions() -> anyhow::Result<()> { let device = Device::Cpu; 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); // Test different hidden dimensions for hidden_dim in [16, 32, 64, 128, 256] { @@ -350,7 +351,7 @@ mod tests { fn test_residual_numerical_stability() -> anyhow::Result<()> { let device = Device::Cpu; 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 config = ResidualConfig { hidden_dim: 32, diff --git a/crates/ml/src/dqn/rmsnorm.rs b/crates/ml/src/dqn/rmsnorm.rs index b5c982f41..2bedc5617 100644 --- a/crates/ml/src/dqn/rmsnorm.rs +++ b/crates/ml/src/dqn/rmsnorm.rs @@ -228,15 +228,16 @@ impl LayerNorm { #[cfg(test)] mod tests { use super::*; - use candle_core::{Device, DType}; + use candle_core::Device; use candle_nn::VarMap; use std::time::Instant; + use crate::dqn::mixed_precision::training_dtype; #[test] fn test_rmsnorm_creation() -> anyhow::Result<()> { let device = Device::Cpu; let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let dim = 128; let rmsnorm = RMSNorm::new_default(vs.pp("rmsnorm"), dim)?; @@ -251,7 +252,7 @@ mod tests { fn test_layernorm_creation() -> anyhow::Result<()> { let device = Device::Cpu; let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let dim = 128; let layernorm = LayerNorm::new_default(vs.pp("layernorm"), dim)?; @@ -266,7 +267,7 @@ mod tests { fn test_rmsnorm_forward() -> anyhow::Result<()> { let device = Device::Cpu; let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let batch_size = 4; let dim = 128; @@ -306,7 +307,7 @@ mod tests { fn test_layernorm_forward() -> anyhow::Result<()> { let device = Device::Cpu; let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let batch_size = 4; let dim = 128; @@ -350,12 +351,12 @@ mod tests { // Setup RMSNorm let rmsnorm_varmap = VarMap::new(); - let rmsnorm_vs = VarBuilder::from_varmap(&rmsnorm_varmap, DType::F32, &device); + let rmsnorm_vs = VarBuilder::from_varmap(&rmsnorm_varmap, training_dtype(&device), &device); let rmsnorm = RMSNorm::new_default(rmsnorm_vs.pp("rmsnorm"), dim)?; // Setup LayerNorm let layernorm_varmap = VarMap::new(); - let layernorm_vs = VarBuilder::from_varmap(&layernorm_varmap, DType::F32, &device); + let layernorm_vs = VarBuilder::from_varmap(&layernorm_varmap, training_dtype(&device), &device); let layernorm = LayerNorm::new_default(layernorm_vs.pp("layernorm"), dim)?; // Create random input @@ -403,11 +404,11 @@ mod tests { // Setup both norms let rmsnorm_varmap = VarMap::new(); - let rmsnorm_vs = VarBuilder::from_varmap(&rmsnorm_varmap, DType::F32, &device); + let rmsnorm_vs = VarBuilder::from_varmap(&rmsnorm_varmap, training_dtype(&device), &device); let rmsnorm = RMSNorm::new_default(rmsnorm_vs.pp("rmsnorm"), dim)?; let layernorm_varmap = VarMap::new(); - let layernorm_vs = VarBuilder::from_varmap(&layernorm_varmap, DType::F32, &device); + let layernorm_vs = VarBuilder::from_varmap(&layernorm_varmap, training_dtype(&device), &device); let layernorm = LayerNorm::new_default(layernorm_vs.pp("layernorm"), dim)?; // Create random input @@ -455,7 +456,7 @@ mod tests { fn test_rmsnorm_3d_input() -> anyhow::Result<()> { let device = Device::Cpu; let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let batch_size = 4; let seq_len = 16; @@ -481,7 +482,7 @@ mod tests { fn test_layernorm_3d_input() -> anyhow::Result<()> { let device = Device::Cpu; let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let batch_size = 4; let seq_len = 16; diff --git a/crates/ml/src/dqn/spectral_norm.rs b/crates/ml/src/dqn/spectral_norm.rs index 09c77e4f4..3f7767988 100644 --- a/crates/ml/src/dqn/spectral_norm.rs +++ b/crates/ml/src/dqn/spectral_norm.rs @@ -232,12 +232,13 @@ impl Module for SpectralNorm { mod tests { use super::*; use candle_nn::VarMap; + use crate::dqn::mixed_precision::training_dtype; #[test] fn test_spectral_norm_creation() -> anyhow::Result<()> { let device = Device::Cpu; let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let config = SpectralNormConfig::default(); let spectral_norm = SpectralNorm::new(10, 5, vs.pp("test"), config)?; @@ -251,7 +252,7 @@ mod tests { fn test_spectral_norm_computation() -> anyhow::Result<()> { let device = Device::Cpu; let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let config = SpectralNormConfig::default(); let mut spectral_norm = SpectralNorm::new(10, 5, vs.pp("test"), config)?; @@ -272,7 +273,7 @@ mod tests { fn test_spectral_norm_bounds_lipschitz() -> anyhow::Result<()> { let device = Device::Cpu; let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let config = SpectralNormConfig { n_power_iterations: 5, // More iterations for accuracy @@ -304,7 +305,7 @@ mod tests { fn test_power_iteration_convergence() -> anyhow::Result<()> { let device = Device::Cpu; let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); // Test with different iteration counts for n_iters in [1, 2, 5] { @@ -327,7 +328,7 @@ mod tests { fn test_singular_vector_reset() -> anyhow::Result<()> { let device = Device::Cpu; let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let config = SpectralNormConfig::default(); let mut spectral_norm = SpectralNorm::new(10, 5, vs.pp("test"), config)?; @@ -352,7 +353,7 @@ mod tests { fn test_forward_pass() -> anyhow::Result<()> { let device = Device::Cpu; let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let config = SpectralNormConfig::default(); let spectral_norm = SpectralNorm::new(10, 5, vs.pp("test"), config)?; @@ -373,7 +374,7 @@ mod tests { fn test_prevents_weight_explosion() -> anyhow::Result<()> { let device = Device::Cpu; let varmap = VarMap::new(); - let vs = VarBuilder::from_varmap(&varmap, DType::F32, &device); + let vs = VarBuilder::from_varmap(&varmap, training_dtype(&device), &device); let config = SpectralNormConfig::default(); let mut spectral_norm = SpectralNorm::new(10, 5, vs.pp("test"), config)?;