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:
jgrusewski
2026-03-03 16:20:15 +01:00
parent 4e0225e090
commit 6febee532f
16 changed files with 78 additions and 64 deletions

View File

@@ -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();

View File

@@ -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)?;

View File

@@ -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

View File

@@ -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();

View File

@@ -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`

View File

@@ -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();

View File

@@ -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(

View File

@@ -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)))?;

View File

@@ -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)?;

View File

@@ -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(

View File

@@ -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,

View File

@@ -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

View File

@@ -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;

View File

@@ -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,

View File

@@ -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;

View File

@@ -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)?;