fix(ml): keep PPO and ensemble pipelines in BF16 on Ampere+ GPUs

Same class of bug as the DQN fix (4c88498b): network outputs were
being cast to F32 mid-pipeline, defeating tensor-core acceleration
on H100/L40S. Now the full training loop stays in training_dtype()
(BF16 on CUDA Ampere+, F32 on CPU), with F32 casts only at scalar
extraction boundaries (to_scalar, to_vec1).

Files fixed:
- ppo.rs: Actor/Critic forward, act_with_log_prob, compute_losses,
  update_mlp, LSTM recurrent loop, predict method
- lstm_networks.rs: removed F32 output casts from both networks
- continuous_ppo.rs: one_tensor + scalar extractions
- hidden_state_manager.rs: zeros/ones use training_dtype()
- flow_policy/mod.rs: log_det accumulators + dummy log_std
- ensemble/adapters/ppo.rs + dqn.rs: F32 cast at extraction

2704 tests pass, 0 clippy warnings.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-04 13:51:23 +01:00
parent 4c88498b8b
commit 2e40a19d6f
7 changed files with 65 additions and 54 deletions

View File

@@ -89,6 +89,8 @@ impl ModelInferenceAdapter for DqnInferenceAdapter {
.squeeze(0)
.map_err(|e| MLError::ModelError(format!("Failed to squeeze Q-values: {e}")))?;
let q_vec_f32: Vec<f32> = q_squeezed
.to_dtype(candle_core::DType::F32)
.map_err(|e| MLError::ModelError(format!("Failed to cast Q-values to F32: {e}")))?
.to_vec1()
.map_err(|e| MLError::ModelError(format!("Failed to extract Q-values: {e}")))?;
let q_vec: Vec<f64> = q_vec_f32.iter().map(|&v| v as f64).collect();

View File

@@ -96,6 +96,8 @@ impl ModelInferenceAdapter for PpoInferenceAdapter {
.squeeze(0)
.map_err(|e| MLError::ModelError(format!("Failed to squeeze probabilities: {e}")))?;
let probs: Vec<f32> = probs_squeezed
.to_dtype(candle_core::DType::F32)
.map_err(|e| MLError::ModelError(format!("Failed to cast probabilities to F32: {e}")))?
.to_vec1()
.map_err(|e| MLError::ModelError(format!("Failed to extract probabilities: {e}")))?;

View File

@@ -4,6 +4,7 @@
//! action spaces, using Gaussian policies for position sizing.
use candle_core::{DType, Device, Tensor};
use crate::dqn::mixed_precision::training_dtype;
use candle_nn::Optimizer; // Required for Adam::new and backward_step methods
use candle_optimisers::adam::Adam;
use candle_optimisers::adam::ParamsAdam;
@@ -521,10 +522,14 @@ impl ContinuousPPO {
})?;
}
total_policy_loss += policy_loss.to_scalar::<f32>().map_err(|e| {
total_policy_loss += policy_loss.to_dtype(DType::F32).map_err(|e| {
MLError::TrainingError(format!("Failed to cast policy loss to F32: {}", e))
})?.to_scalar::<f32>().map_err(|e| {
MLError::TrainingError(format!("Failed to extract policy loss: {}", e))
})?;
total_value_loss += value_loss.to_scalar::<f32>().map_err(|e| {
total_value_loss += value_loss.to_dtype(DType::F32).map_err(|e| {
MLError::TrainingError(format!("Failed to cast value loss to F32: {}", e))
})?.to_scalar::<f32>().map_err(|e| {
MLError::TrainingError(format!("Failed to extract value loss: {}", e))
})?;
num_updates += 1;
@@ -565,7 +570,8 @@ impl ContinuousPPO {
)
.map_err(|e| MLError::TrainingError(format!("Failed to create clip tensor: {}", e)))?;
let one_tensor = Tensor::ones(batch.advantages.dims(), DType::F32, self.actor.device())?;
let dtype = training_dtype(self.actor.device());
let one_tensor = Tensor::ones(batch.advantages.dims(), dtype, self.actor.device())?;
let clip_min = (&one_tensor - &clip_epsilon_tensor)?;
let clip_max = (&one_tensor + &clip_epsilon_tensor)?;

View File

@@ -3,7 +3,7 @@ mod flow_matching;
use coupling_layer::AffineCouplingLayer;
use candle_core::{DType, Device, Tensor};
use candle_core::{Device, Tensor};
use candle_nn::{linear, Linear, Module, VarBuilder, VarMap};
use rand::thread_rng;
use rand_distr::{Distribution, Normal};
@@ -427,7 +427,7 @@ impl FlowPolicy {
fn flow_forward(&self, z: &Tensor, ctx: &Tensor) -> Result<(Tensor, Tensor), MLError> {
let mut x = z.clone();
let batch_size = z.dims()[0];
let mut log_det_acc = Tensor::zeros(batch_size, DType::F32, &self.device)
let mut log_det_acc = Tensor::zeros(batch_size, training_dtype(&self.device), &self.device)
.map_err(|e| MLError::TensorOperationError(format!("Log det init failed: {}", e)))?;
for layer in &self.layers {
@@ -452,7 +452,7 @@ impl FlowPolicy {
fn flow_inverse(&self, y: &Tensor, ctx: &Tensor) -> Result<(Tensor, Tensor), MLError> {
let mut x = y.clone();
let batch_size = y.dims()[0];
let mut log_det_acc = Tensor::zeros(batch_size, DType::F32, &self.device)
let mut log_det_acc = Tensor::zeros(batch_size, training_dtype(&self.device), &self.device)
.map_err(|e| MLError::TensorOperationError(format!("Log det init failed: {}", e)))?;
// Reverse order of layers for inverse
@@ -524,7 +524,7 @@ impl FlowPolicy {
/// Normalizing flows don't have a fixed log_std parameter like Gaussian policies.
/// Returns zeros for API compatibility with existing PPO code.
pub fn get_current_log_std(&self) -> Result<Tensor, MLError> {
Tensor::zeros(self.config.action_dim, DType::F32, &self.device)
Tensor::zeros(self.config.action_dim, training_dtype(&self.device), &self.device)
.map_err(|e| MLError::TensorOperationError(format!("Log std creation failed: {}", e)))
}

View File

@@ -3,9 +3,12 @@
//! Manages LSTM hidden states (h_t, c_t) across timesteps and episodes.
//! States persist within episodes but reset at episode boundaries.
use candle_core::{DType, Device, Tensor};
use candle_core::{Device, Tensor};
#[cfg(test)]
use candle_core::DType;
use std::fmt;
use crate::MLError;
use crate::dqn::mixed_precision::training_dtype;
/// Manages LSTM hidden and cell states for policy and value networks
pub struct HiddenStateManager {
@@ -40,7 +43,7 @@ impl HiddenStateManager {
device: &Device,
) -> Result<Self, MLError> {
let shape = &[num_layers, batch_size, hidden_dim];
let zeros = Tensor::zeros(shape, DType::F32, device)
let zeros = Tensor::zeros(shape, training_dtype(device), device)
.map_err(|e| MLError::TensorOperationError(format!("Failed to create zero tensor: {}", e)))?;
Ok(Self {
@@ -137,8 +140,9 @@ impl HiddenStateManager {
// Convert done_mask to float and expand to match state dimensions
// done_mask: [batch_size] -> [1, batch_size, 1]
let dtype = training_dtype(&self.device);
let done_float = done_mask
.to_dtype(DType::F32)
.to_dtype(dtype)
.map_err(|e| MLError::TensorOperationError(format!("Failed to convert done mask to float: {}", e)))?;
let done_expanded = done_float
@@ -153,7 +157,7 @@ impl HiddenStateManager {
.map_err(|e| MLError::TensorOperationError(format!("Failed to broadcast done mask: {}", e)))?;
// Create keep_mask = 1 - done_mask (keep states where episode continues)
let ones = Tensor::ones(&[self.num_layers, self.batch_size, self.hidden_dim], DType::F32, &self.device)
let ones = Tensor::ones(&[self.num_layers, self.batch_size, self.hidden_dim], dtype, &self.device)
.map_err(|e| MLError::TensorOperationError(format!("Failed to create ones tensor: {}", e)))?;
let keep_mask = ones
@@ -183,7 +187,7 @@ impl HiddenStateManager {
/// Reset all states to zeros
pub fn reset_all(&mut self) -> Result<(), MLError> {
let shape = &[self.num_layers, self.batch_size, self.hidden_dim];
let zeros = Tensor::zeros(shape, DType::F32, &self.device)
let zeros = Tensor::zeros(shape, training_dtype(&self.device), &self.device)
.map_err(|e| MLError::TensorOperationError(format!("Failed to create zero tensor: {}", e)))?;
self.policy_hidden = zeros.clone();

View File

@@ -162,14 +162,6 @@ impl LSTMPolicyNetwork {
.forward(&x)
.map_err(|e| MLError::ModelError(format!("Output layer forward failed: {}", e)))?;
// Cast outputs back to F32 for API compatibility
let logits = logits.to_dtype(candle_core::DType::F32)
.map_err(|e| MLError::ModelError(format!("Logits dtype cast failed: {}", e)))?;
let new_h = new_h.to_dtype(candle_core::DType::F32)
.map_err(|e| MLError::ModelError(format!("Hidden state dtype cast failed: {}", e)))?;
let new_c = new_c.to_dtype(candle_core::DType::F32)
.map_err(|e| MLError::ModelError(format!("Cell state dtype cast failed: {}", e)))?;
Ok((logits, new_h, new_c))
}
@@ -391,14 +383,6 @@ impl LSTMValueNetwork {
.squeeze(1)
.map_err(|e| MLError::ModelError(format!("Failed to squeeze value output: {}", e)))?;
// Cast outputs back to F32 for API compatibility
let value = value.to_dtype(candle_core::DType::F32)
.map_err(|e| MLError::ModelError(format!("Value dtype cast failed: {}", e)))?;
let new_h = new_h.to_dtype(candle_core::DType::F32)
.map_err(|e| MLError::ModelError(format!("Hidden state dtype cast failed: {}", e)))?;
let new_c = new_c.to_dtype(candle_core::DType::F32)
.map_err(|e| MLError::ModelError(format!("Cell state dtype cast failed: {}", e)))?;
Ok((value, new_h, new_c))
}

View File

@@ -464,13 +464,6 @@ impl PolicyNetwork {
}
}
// Always cast output back to F32 for API compatibility
if x.dtype() != DType::F32 {
x = x.to_dtype(DType::F32).map_err(|e| {
MLError::ModelError(format!("Output dtype cast to F32 failed: {}", e))
})?;
}
Ok(x)
}
@@ -716,13 +709,6 @@ impl ValueNetwork {
// Squeeze the last dimension (from [batch, 1] to [batch])
x = x.squeeze(1)?;
// Always cast output back to F32 for API compatibility
if x.dtype() != DType::F32 {
x = x.to_dtype(DType::F32).map_err(|e| {
MLError::ModelError(format!("Output dtype cast to F32 failed: {}", e))
})?;
}
Ok(x)
}
@@ -950,11 +936,13 @@ impl PPO {
// Get action and log-probability from policy
let (action, log_prob) = self.actor.sample_action(&state_tensor)?;
// Get value estimate
// Get value estimate — cast to F32 for scalar extraction
let value_tensor = self.critic.forward(&state_tensor)?;
let value = value_tensor
.get(0)
.map_err(|e| MLError::ModelError(format!("Failed to get value element: {}", e)))?
.to_dtype(DType::F32)
.map_err(|e| MLError::ModelError(format!("Failed to cast value to F32: {}", e)))?
.to_scalar::<f32>()
.map_err(|e| MLError::ModelError(format!("Failed to extract value: {}", e)))?;
@@ -1080,7 +1068,9 @@ impl PPO {
// Only compute value loss (no policy loss)
let value_loss = self.compute_value_loss(&mini_tensors)?;
let value_loss_scalar = value_loss.to_scalar::<f32>().map_err(|e| {
let value_loss_scalar = value_loss.to_dtype(DType::F32).map_err(|e| {
MLError::TrainingError(format!("Failed to cast value loss to F32: {}", e))
})?.to_scalar::<f32>().map_err(|e| {
MLError::TrainingError(format!("Failed to extract value loss: {}", e))
})?;
@@ -1204,12 +1194,17 @@ impl PPO {
let value_loss = self.compute_value_loss(&mini_tensors)?;
// Extract scalar values for NaN check and loss tracking.
// Cast to F32 at boundary — loss tensors may be BF16 on Ampere+ GPUs.
// PPO has ~40-80 mini-batches per update (low count), so per-batch
// GPU sync overhead is acceptable unlike DQN's thousands of steps.
let policy_loss_scalar = policy_loss.to_scalar::<f32>().map_err(|e| {
let policy_loss_scalar = policy_loss.to_dtype(DType::F32).map_err(|e| {
MLError::TrainingError(format!("Failed to cast policy loss to F32: {}", e))
})?.to_scalar::<f32>().map_err(|e| {
MLError::TrainingError(format!("Failed to extract policy loss: {}", e))
})?;
let value_loss_scalar = value_loss.to_scalar::<f32>().map_err(|e| {
let value_loss_scalar = value_loss.to_dtype(DType::F32).map_err(|e| {
MLError::TrainingError(format!("Failed to cast value loss to F32: {}", e))
})?.to_scalar::<f32>().map_err(|e| {
MLError::TrainingError(format!("Failed to extract value loss: {}", e))
})?;
@@ -1506,6 +1501,10 @@ impl PPO {
let log_prob = log_probs_dist.gather(&action_tensor, 1)?
.squeeze(1)?
.get(0)?
.to_dtype(DType::F32)
.map_err(|e| MLError::TensorOperationError(
format!("Failed to cast log prob to F32: {}", e)
))?
.to_scalar::<f32>()
.map_err(|e| MLError::TensorOperationError(
format!("Failed to extract log prob: {}", e)
@@ -1518,12 +1517,20 @@ impl PPO {
.sum(candle_core::D::Minus1)?;
let step_entropy = TensorOps::negate(&step_entropy_inner)?
.get(0)?
.to_dtype(DType::F32)
.map_err(|e| MLError::TensorOperationError(
format!("Failed to cast step entropy to F32: {}", e)
))?
.to_scalar::<f32>()
.map_err(|e| MLError::TensorOperationError(
format!("Failed to extract step entropy: {}", e)
))?;
let value_scalar = value.get(0)?
.to_dtype(DType::F32)
.map_err(|e| MLError::TensorOperationError(
format!("Failed to cast value to F32: {}", e)
))?
.to_scalar::<f32>()
.map_err(|e| MLError::TensorOperationError(
format!("Failed to extract value: {}", e)
@@ -1580,7 +1587,8 @@ impl PPO {
device,
)?;
let one_tensor = Tensor::ones((seq_len,), DType::F32, device)?;
let dtype = training_dtype(device);
let one_tensor = Tensor::ones((seq_len,), dtype, device)?;
let clip_min = (&one_tensor - &clip_epsilon_tensor)?;
let clip_max = if let Some(eps_high) = self.config.clip_epsilon_high {
let clip_high_tensor = Tensor::from_vec(
@@ -1631,11 +1639,15 @@ impl PPO {
// Track mean log probability for adaptive entropy update
last_mean_log_pi = Some(seq_new_log_probs.mean_all()?);
// Extract scalar values for NaN check
let policy_loss_scalar = policy_loss.to_scalar::<f32>().map_err(|e| {
// Extract scalar values for NaN check — cast to F32 at extraction boundary
let policy_loss_scalar = policy_loss.to_dtype(DType::F32).map_err(|e| {
MLError::TrainingError(format!("Failed to cast policy loss to F32: {}", e))
})?.to_scalar::<f32>().map_err(|e| {
MLError::TrainingError(format!("Failed to extract policy loss: {}", e))
})?;
let value_loss_scalar = scaled_value_loss.to_scalar::<f32>().map_err(|e| {
let value_loss_scalar = scaled_value_loss.to_dtype(DType::F32).map_err(|e| {
MLError::TrainingError(format!("Failed to cast value loss to F32: {}", e))
})?.to_scalar::<f32>().map_err(|e| {
MLError::TrainingError(format!("Failed to extract value loss: {}", e))
})?;
@@ -1758,8 +1770,8 @@ impl PPO {
let policy_loss = self.compute_policy_loss(&batch_tensors)?;
let value_loss = self.compute_value_loss(&batch_tensors)?;
let policy_loss_scalar = policy_loss.to_scalar::<f32>()?;
let value_loss_scalar = value_loss.to_scalar::<f32>()?;
let policy_loss_scalar = policy_loss.to_dtype(DType::F32)?.to_scalar::<f32>()?;
let value_loss_scalar = value_loss.to_dtype(DType::F32)?.to_scalar::<f32>()?;
Ok((policy_loss_scalar, value_loss_scalar))
}
@@ -1781,7 +1793,8 @@ impl PPO {
)
.map_err(|e| MLError::TrainingError(format!("Failed to create clip tensor: {}", e)))?;
let one_tensor = Tensor::ones(batch.advantages.dims(), DType::F32, self.actor.device())?;
let dtype = training_dtype(self.actor.device());
let one_tensor = Tensor::ones(batch.advantages.dims(), dtype, self.actor.device())?;
let clip_min = (&one_tensor - &clip_epsilon_tensor)?;
let clip_max = if let Some(eps_high) = self.config.clip_epsilon_high {
let clip_high_tensor = Tensor::from_vec(
@@ -2297,7 +2310,7 @@ impl PPO {
self.actor.device(),
)?;
let probs_tensor = self.actor.action_probabilities(&state_tensor)?;
let probs = probs_tensor.flatten_all()?.to_vec1::<f32>()?;
let probs = probs_tensor.to_dtype(DType::F32)?.flatten_all()?.to_vec1::<f32>()?;
Ok(probs)
}