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:
@@ -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();
|
||||
|
||||
@@ -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)?;
|
||||
|
||||
|
||||
Reference in New Issue
Block a user