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>
227 lines
7.4 KiB
Rust
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(())
|
|
}
|
|
}
|