diff --git a/crates/ml-dqn/src/branching.rs b/crates/ml-dqn/src/branching.rs index 370eb75d9..6bf0c977d 100644 --- a/crates/ml-dqn/src/branching.rs +++ b/crates/ml-dqn/src/branching.rs @@ -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(); diff --git a/crates/ml-dqn/src/noisy_layers.rs b/crates/ml-dqn/src/noisy_layers.rs index 2dc2e2099..c497953e0 100644 --- a/crates/ml-dqn/src/noisy_layers.rs +++ b/crates/ml-dqn/src/noisy_layers.rs @@ -64,8 +64,8 @@ impl NoisyLinear { ) -> Result { 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)?;