Files
foxhunt/crates/ml-ppo/src/hidden_state_manager.rs
jgrusewski 22004a7368 refactor(cuda): eliminate candle from ml-core, ml-ppo, and 4 thin crates
Hard refactor — no shims, no compat layers. Candle removed from Cargo.toml
and all source files in 6 crates:

- ml-core: MlDevice enum, checkpoint.rs (safetensors direct), cudarc imports
  fixed from candle re-export to direct, AdamWConfig lr_decay, cuda_compat
  gutted. Net -7,341 lines.
- ml-ppo: All 16 files rewritten. LSTM→CudaLSTM, VarMap→GpuVarStore,
  PPOAgent 2306→700 lines, checkpoint→binary format.
- ml-ensemble: GPU-resident sigmoid via custom CUDA kernel.
- ml-explainability: Integrated gradients via GPU finite-difference kernels.
- ml-labeling: Device→MlDevice.
- ml-hyperopt: Cargo.toml only.

Remaining: ml-dqn (24 files), ml-supervised (4 files), ml crate (104 files).

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-03-17 22:27:56 +01:00

227 lines
7.4 KiB
Rust

//! Hidden State Management for Recurrent PPO
//!
//! Manages LSTM hidden states (`h_t`, `c_t`) across timesteps and episodes.
//! States persist within episodes but reset at episode boundaries.
//! All state storage is host-side (`Vec<f32>`); GPU upload happens in the LSTM forward pass.
use std::fmt;
use ml_core::MLError;
/// Manages LSTM hidden and cell states for policy and value networks
pub struct HiddenStateManager {
/// Policy network hidden state (flat: `num_layers * batch_size * hidden_dim`)
policy_hidden: Vec<f32>,
/// Policy network cell state
policy_cell: Vec<f32>,
/// Value network hidden state
value_hidden: Vec<f32>,
/// Value network cell state
value_cell: Vec<f32>,
/// Dimensions for creating new tensors
num_layers: usize,
batch_size: usize,
hidden_dim: usize,
}
impl HiddenStateManager {
/// Create a new hidden state manager with all states initialized to zeros
pub fn new(
num_layers: usize,
batch_size: usize,
hidden_dim: usize,
_device: &(), // Kept for API compat; states are host-side
) -> Result<Self, MLError> {
let total = num_layers * batch_size * hidden_dim;
Ok(Self {
policy_hidden: vec![0.0; total],
policy_cell: vec![0.0; total],
value_hidden: vec![0.0; total],
value_cell: vec![0.0; total],
num_layers,
batch_size,
hidden_dim,
})
}
/// Create with default device arg (unit type)
pub fn with_defaults(
num_layers: usize,
batch_size: usize,
hidden_dim: usize,
) -> Result<Self, MLError> {
Self::new(num_layers, batch_size, hidden_dim, &())
}
/// Get policy network states (hidden, cell) as flat f32 slices
pub fn get_policy_state(&self) -> (&[f32], &[f32]) {
(&self.policy_hidden, &self.policy_cell)
}
/// Get value network states (hidden, cell) as flat f32 slices
pub fn get_value_state(&self) -> (&[f32], &[f32]) {
(&self.value_hidden, &self.value_cell)
}
/// Update policy network states
pub fn update_policy_state(&mut self, hidden: Vec<f32>, cell: Vec<f32>) -> Result<(), MLError> {
let expected = self.num_layers * self.batch_size * self.hidden_dim;
if hidden.len() != expected {
return Err(MLError::InvalidInput(format!(
"Invalid hidden state length. Expected {}, got {}",
expected, hidden.len()
)));
}
if cell.len() != expected {
return Err(MLError::InvalidInput(format!(
"Invalid cell state length. Expected {}, got {}",
expected, cell.len()
)));
}
self.policy_hidden = hidden;
self.policy_cell = cell;
Ok(())
}
/// Update value network states
pub fn update_value_state(&mut self, hidden: Vec<f32>, cell: Vec<f32>) -> Result<(), MLError> {
let expected = self.num_layers * self.batch_size * self.hidden_dim;
if hidden.len() != expected {
return Err(MLError::InvalidInput(format!(
"Invalid hidden state length. Expected {}, got {}",
expected, hidden.len()
)));
}
if cell.len() != expected {
return Err(MLError::InvalidInput(format!(
"Invalid cell state length. Expected {}, got {}",
expected, cell.len()
)));
}
self.value_hidden = hidden;
self.value_cell = cell;
Ok(())
}
/// Reset states for done environments
pub fn reset_on_done(&mut self, done_mask: &[bool]) -> Result<(), MLError> {
if done_mask.len() != self.batch_size {
return Err(MLError::InvalidInput(format!(
"Invalid done mask length. Expected {}, got {}",
self.batch_size,
done_mask.len()
)));
}
for (batch_idx, &done) in done_mask.iter().enumerate() {
if done {
for layer in 0..self.num_layers {
let offset = layer * self.batch_size * self.hidden_dim
+ batch_idx * self.hidden_dim;
for h in 0..self.hidden_dim {
let idx = offset + h;
if let Some(v) = self.policy_hidden.get_mut(idx) {
*v = 0.0;
}
if let Some(v) = self.policy_cell.get_mut(idx) {
*v = 0.0;
}
if let Some(v) = self.value_hidden.get_mut(idx) {
*v = 0.0;
}
if let Some(v) = self.value_cell.get_mut(idx) {
*v = 0.0;
}
}
}
}
}
Ok(())
}
/// Reset all states to zeros
pub fn reset_all(&mut self) -> Result<(), MLError> {
let total = self.num_layers * self.batch_size * self.hidden_dim;
self.policy_hidden = vec![0.0; total];
self.policy_cell = vec![0.0; total];
self.value_hidden = vec![0.0; total];
self.value_cell = vec![0.0; total];
Ok(())
}
/// Get dimensions
pub const fn num_layers(&self) -> usize { self.num_layers }
pub const fn batch_size(&self) -> usize { self.batch_size }
pub const fn hidden_dim(&self) -> usize { self.hidden_dim }
}
impl fmt::Debug for HiddenStateManager {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("HiddenStateManager")
.field("num_layers", &self.num_layers)
.field("batch_size", &self.batch_size)
.field("hidden_dim", &self.hidden_dim)
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_new_creates_zero_states() -> Result<(), MLError> {
let manager = HiddenStateManager::with_defaults(2, 4, 64)?;
let (ph, pc) = manager.get_policy_state();
let (vh, vc) = manager.get_value_state();
assert_eq!(ph.len(), 2 * 4 * 64);
assert_eq!(pc.len(), 2 * 4 * 64);
assert_eq!(vh.len(), 2 * 4 * 64);
assert_eq!(vc.len(), 2 * 4 * 64);
assert!(ph.iter().all(|&v| v == 0.0));
Ok(())
}
#[test]
fn test_state_updates() -> Result<(), MLError> {
let mut manager = HiddenStateManager::with_defaults(1, 2, 3)?;
let new_h = vec![1.0; 6];
let new_c = vec![1.0; 6];
manager.update_policy_state(new_h, new_c)?;
let (ph, pc) = manager.get_policy_state();
let ph_sum: f32 = ph.iter().sum();
let pc_sum: f32 = pc.iter().sum();
assert!((ph_sum - 6.0).abs() < 0.01);
assert!((pc_sum - 6.0).abs() < 0.01);
Ok(())
}
#[test]
fn test_reset_all() -> Result<(), MLError> {
let mut manager = HiddenStateManager::with_defaults(1, 2, 3)?;
let ones = vec![1.0; 6];
manager.update_policy_state(ones.clone(), ones.clone())?;
manager.update_value_state(ones.clone(), ones)?;
manager.reset_all()?;
let (ph, pc) = manager.get_policy_state();
let (vh, vc) = manager.get_value_state();
assert!(ph.iter().all(|&v| v == 0.0));
assert!(pc.iter().all(|&v| v == 0.0));
assert!(vh.iter().all(|&v| v == 0.0));
assert!(vc.iter().all(|&v| v == 0.0));
Ok(())
}
}