feat(ml): BF16 VarBuilder for all DQN networks and layers
Replace DType::F32 with training_dtype(&device) in all VarBuilder::from_varmap calls across 16 DQN files (~55 call sites). This enables automatic BF16 weight initialization on Ampere+ GPUs while keeping F32 on CPU and older hardware. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -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();
|
||||
|
||||
@@ -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)?;
|
||||
|
||||
|
||||
@@ -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<Self, MLError> {
|
||||
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
|
||||
|
||||
@@ -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<Self, MLError> {
|
||||
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();
|
||||
|
||||
@@ -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`
|
||||
|
||||
@@ -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<Self, MLError> {
|
||||
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();
|
||||
|
||||
@@ -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<Self, MLError> {
|
||||
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(
|
||||
|
||||
@@ -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)))?;
|
||||
|
||||
|
||||
@@ -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)?;
|
||||
|
||||
|
||||
@@ -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<Self, MLError> {
|
||||
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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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)?;
|
||||
|
||||
Reference in New Issue
Block a user