fix(dqn): init weights as F32, not BF16 — fixes ensure_f32 Var::set crash

VarBuilder and NoisyLinear were creating BF16 weights, then ensure_f32
tried Var::set() which rejects dtype changes. Fix: create F32 from the
start. BF16 mirrors are managed separately by GpuDqnTrainer.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-17 15:34:19 +01:00
parent 0e2f82ab54
commit 1627ca57a1
2 changed files with 17 additions and 15 deletions

View File

@@ -297,7 +297,9 @@ impl BranchingDuelingQNetwork {
}
let vars = VarMap::new();
let vb = VarBuilder::from_varmap(&vars, candle_core::DType::BF16, &device);
// F32 weights: the fused CUDA trainer's Adam kernel and gpu_weights.rs
// fast-path extraction require F32. BF16 mirrors are maintained by GpuDqnTrainer.
let vb = VarBuilder::from_varmap(&vars, candle_core::DType::F32, &device);
// Shared encoder (always standard Linear -- noise only in heads)
let mut shared_layers = Vec::new();

View File

@@ -64,8 +64,8 @@ impl NoisyLinear {
) -> Result<Self, MLError> {
let device = vb.device().clone();
// All init in F32, then cast — Candle's Init::Uniform CUDA kernel lacks BF16 PTX.
let dtype = candle_core::DType::BF16;
// F32 weights: fused CUDA trainer operates on F32, BF16 mirrors managed separately.
let dtype = candle_core::DType::F32;
// Initialize μ_w ~ U(-1/√in, 1/√in) (Rainbow DQN standard)
let mu_range = 1.0 / (in_features as f64).sqrt();
@@ -384,7 +384,7 @@ mod tests {
fn test_noisy_linear_creation() -> Result<(), MLError> {
let device = Device::new_cuda(0).expect("CUDA required");
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device);
let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::F32, &device);
let _layer = NoisyLinear::new(64, 32, vb, 0.5)?;
Ok(())
@@ -394,7 +394,7 @@ mod tests {
fn test_noisy_linear_forward() -> Result<(), MLError> {
let device = Device::new_cuda(0).expect("CUDA required");
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device);
let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::F32, &device);
let mut layer = NoisyLinear::new(64, 32, vb, 0.5)?;
layer.reset_noise()?; // Resample noise before forward
@@ -402,7 +402,7 @@ mod tests {
// Create dummy input
let input = Tensor::randn(0.0_f32, 1.0_f32, (4, 64), &device)
.map_err(|e| MLError::ModelError(format!("Failed to create input: {}", e)))?
.to_dtype(candle_core::DType::BF16)
.to_dtype(candle_core::DType::F32)
.map_err(|e| MLError::ModelError(format!("Failed to cast input: {}", e)))?;
// Forward pass
@@ -418,12 +418,12 @@ mod tests {
fn test_noise_reset() -> Result<(), MLError> {
let device = Device::new_cuda(0).expect("CUDA required");
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device);
let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::F32, &device);
let mut layer = NoisyLinear::new(64, 32, vb, 0.5)?;
let input = Tensor::randn(0.0_f32, 1.0_f32, (4, 64), &device)
.map_err(|e| MLError::ModelError(format!("Failed to create input: {}", e)))?
.to_dtype(candle_core::DType::BF16)
.to_dtype(candle_core::DType::F32)
.map_err(|e| MLError::ModelError(format!("Failed to cast input: {}", e)))?;
// First forward pass
@@ -467,12 +467,12 @@ mod tests {
fn test_disable_noise() -> Result<(), MLError> {
let device = Device::new_cuda(0).expect("CUDA required");
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device);
let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::F32, &device);
let mut layer = NoisyLinear::new(64, 32, vb, 0.5)?;
let input = Tensor::randn(0.0_f32, 1.0_f32, (4, 64), &device)
.map_err(|e| MLError::ModelError(format!("Failed to create input: {}", e)))?
.to_dtype(candle_core::DType::BF16)
.to_dtype(candle_core::DType::F32)
.map_err(|e| MLError::ModelError(format!("Failed to cast input: {}", e)))?;
// Reset noise for first pass
@@ -511,7 +511,7 @@ mod tests {
fn test_factorized_noise_dimensions() -> Result<(), MLError> {
let device = Device::new_cuda(0).expect("CUDA required");
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device);
let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::F32, &device);
let mut layer = NoisyLinear::new(128, 64, vb, 0.5)?;
layer.reset_noise()?;
@@ -527,12 +527,12 @@ mod tests {
fn test_reset_noise_with_sigma() -> Result<(), MLError> {
let device = Device::new_cuda(0).expect("CUDA required");
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device);
let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::F32, &device);
let mut layer = NoisyLinear::new(64, 32, vb, 0.5)?;
let input = Tensor::randn(0.0_f32, 1.0_f32, (4, 64), &device)
.map_err(|e| MLError::ModelError(format!("Failed to create input: {}", e)))?
.to_dtype(candle_core::DType::BF16)
.to_dtype(candle_core::DType::F32)
.map_err(|e| MLError::ModelError(format!("Failed to cast input: {}", e)))?;
// Test with high sigma (0.6)
@@ -571,7 +571,7 @@ mod tests {
fn test_sigma_scaling_effect() -> Result<(), MLError> {
let device = Device::new_cuda(0).expect("CUDA required");
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::BF16, &device);
let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::F32, &device);
let mut layer = NoisyLinear::new(64, 32, vb, 0.5)?;
@@ -579,7 +579,7 @@ mod tests {
layer.disable_noise()?;
let input = Tensor::randn(0.0_f32, 1.0_f32, (4, 64), &device)
.map_err(|e| MLError::ModelError(format!("Failed to create input: {}", e)))?
.to_dtype(candle_core::DType::BF16)
.to_dtype(candle_core::DType::F32)
.map_err(|e| MLError::ModelError(format!("Failed to cast input: {}", e)))?;
let output_no_noise = layer.forward(&input)?;