16 files: doc backticks (60+), const fn (20+), safety comments (25+), needless borrows, div_ceil, dead fields, split multi-op unsafe blocks, let..else patterns, else-if-without-else, redundant casts. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
208 lines
7.1 KiB
Rust
208 lines
7.1 KiB
Rust
//! Rainbow DQN Agent Implementation
|
|
//!
|
|
//! Complete implementation of Rainbow DQN agent with all 6 components:
|
|
//! 1. Double Q-learning, 2. Dueling Networks, 3. Prioritized Experience Replay,
|
|
//! 4. Multi-step Learning, 5. Distributional RL (C51), 6. Noisy Networks
|
|
//!
|
|
//! NOTE: This module is a cold-path orchestrator. The hot-path forward/backward
|
|
//! runs through fused CUDA kernels (`dqn_experience_kernel.cu`,
|
|
//! `dqn_forward_only_kernel`). This code handles replay buffer management,
|
|
//! target network syncing, and metric tracking.
|
|
|
|
use std::sync::{Arc, Mutex, RwLock};
|
|
|
|
use tracing::info;
|
|
|
|
use super::rainbow_config::{RainbowAgentConfig, RainbowAgentMetrics, TrainingResult};
|
|
use super::{Experience, ReplayBuffer, ReplayBufferConfig};
|
|
use ml_core::MLError;
|
|
|
|
/// Rainbow `DQN` Agent with all 6 components.
|
|
///
|
|
/// Training forward/backward is handled by the fused CUDA DQN trainer.
|
|
/// This struct manages the replay buffer, metrics, and configuration.
|
|
pub struct RainbowAgent {
|
|
config: RainbowAgentConfig,
|
|
|
|
// Experience replay
|
|
replay_buffer: Arc<Mutex<ReplayBuffer>>,
|
|
|
|
// Metrics and state
|
|
metrics: Arc<RwLock<RainbowAgentMetrics>>,
|
|
step_count: Arc<Mutex<u64>>,
|
|
|
|
// Priority replay state
|
|
priority_beta: Arc<Mutex<f64>>,
|
|
}
|
|
|
|
impl RainbowAgent {
|
|
/// Create a new Rainbow `DQN` agent
|
|
pub fn new(config: RainbowAgentConfig) -> Result<Self, MLError> {
|
|
info!("Rainbow Agent initializing (CUDA required for training)");
|
|
|
|
// Create replay buffer
|
|
let buffer_config = ReplayBufferConfig {
|
|
capacity: config.replay_buffer_size,
|
|
batch_size: config.batch_size,
|
|
min_experiences: config.min_replay_size,
|
|
};
|
|
|
|
let replay_buffer = Arc::new(Mutex::new(ReplayBuffer::new(buffer_config)?));
|
|
|
|
// Initialize state
|
|
let metrics = Arc::new(RwLock::new(RainbowAgentMetrics::default()));
|
|
let step_count = Arc::new(Mutex::new(0));
|
|
let priority_beta = Arc::new(Mutex::new(config.priority_beta));
|
|
|
|
Ok(Self {
|
|
config,
|
|
replay_buffer,
|
|
metrics,
|
|
step_count,
|
|
priority_beta,
|
|
})
|
|
}
|
|
|
|
/// Record a step and increment the step counter.
|
|
pub fn record_step(&self) -> Result<u64, MLError> {
|
|
let mut step_count = self
|
|
.step_count
|
|
.lock()
|
|
.map_err(|e| MLError::LockError(format!("Failed to acquire step_count lock: {e}")))?;
|
|
*step_count += 1;
|
|
|
|
let mut metrics = self.metrics.write().map_err(|e| {
|
|
MLError::LockError(format!("Failed to acquire metrics write lock: {e}"))
|
|
})?;
|
|
metrics.total_steps = *step_count;
|
|
|
|
Ok(*step_count)
|
|
}
|
|
|
|
/// Add experience to replay buffer
|
|
pub fn add_experience(&self, experience: Experience) -> Result<(), MLError> {
|
|
let buffer = self.replay_buffer.lock().map_err(|e| {
|
|
MLError::LockError(format!("Failed to acquire replay_buffer lock: {e}"))
|
|
})?;
|
|
buffer.push(experience)?;
|
|
|
|
let mut metrics = self.metrics.write().map_err(|e| {
|
|
MLError::LockError(format!("Failed to acquire metrics write lock: {e}"))
|
|
})?;
|
|
metrics.replay_buffer_size = buffer.size();
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Check if training is possible (enough experiences collected)
|
|
pub fn can_train(&self) -> Result<bool, MLError> {
|
|
let buffer = self.replay_buffer.lock().map_err(|e| {
|
|
MLError::LockError(format!("Failed to acquire replay_buffer lock: {e}"))
|
|
})?;
|
|
Ok(buffer.can_sample() && buffer.size() >= self.config.min_replay_size)
|
|
}
|
|
|
|
/// Record a training step result.
|
|
pub fn record_training_result(&self, loss: f64) -> Result<TrainingResult, MLError> {
|
|
// Update priority beta
|
|
{
|
|
let mut beta = self.priority_beta.lock().map_err(|e| {
|
|
MLError::LockError(format!("Failed to acquire priority_beta lock: {e}"))
|
|
})?;
|
|
*beta = (*beta + self.config.priority_beta_increment).min(1.0);
|
|
}
|
|
|
|
// Update metrics
|
|
{
|
|
let mut metrics = self.metrics.write().map_err(|e| {
|
|
MLError::LockError(format!("Failed to acquire metrics write lock: {e}"))
|
|
})?;
|
|
metrics.current_loss = loss;
|
|
let beta = self.priority_beta.lock().map_err(|e| {
|
|
MLError::LockError(format!(
|
|
"Failed to acquire priority_beta lock for metrics update: {e}",
|
|
))
|
|
})?;
|
|
metrics.priority_beta = *beta;
|
|
}
|
|
|
|
Ok(TrainingResult::new(loss))
|
|
}
|
|
|
|
/// Get current metrics
|
|
pub fn metrics(&self) -> RainbowAgentMetrics {
|
|
self.metrics.read().map(|m| m.clone()).unwrap_or_default()
|
|
}
|
|
|
|
/// Get current step count
|
|
pub fn step_count(&self) -> u64 {
|
|
self.step_count.lock().map(|c| *c).unwrap_or(0)
|
|
}
|
|
|
|
/// Check if target network should be updated
|
|
pub fn should_update_target(&self) -> bool {
|
|
let sc = self.step_count();
|
|
sc > 0 && sc % self.config.target_update_freq as u64 == 0
|
|
}
|
|
|
|
/// Reset agent state
|
|
pub fn reset(&self) -> Result<(), MLError> {
|
|
// Reset metrics
|
|
{
|
|
let mut metrics = self.metrics.write().map_err(|e| {
|
|
MLError::LockError(format!("Failed to acquire metrics write lock for reset: {e}"))
|
|
})?;
|
|
*metrics = RainbowAgentMetrics::default();
|
|
}
|
|
|
|
// Reset counters
|
|
{
|
|
let mut step_count = self.step_count.lock().map_err(|e| {
|
|
MLError::LockError(format!("Failed to acquire step_count lock for reset: {e}"))
|
|
})?;
|
|
*step_count = 0;
|
|
}
|
|
|
|
// Reset replay buffer
|
|
{
|
|
let mut buffer = self.replay_buffer.lock().map_err(|e| {
|
|
MLError::LockError(format!("Failed to acquire replay_buffer lock for reset: {e}"))
|
|
})?;
|
|
let buffer_config = ReplayBufferConfig {
|
|
capacity: self.config.replay_buffer_size,
|
|
batch_size: self.config.batch_size,
|
|
min_experiences: self.config.min_replay_size,
|
|
};
|
|
*buffer = ReplayBuffer::new(buffer_config)?;
|
|
}
|
|
|
|
info!("Rainbow Agent reset completed");
|
|
Ok(())
|
|
}
|
|
|
|
/// Get the replay buffer (for external sampling by the fused CUDA trainer).
|
|
pub const fn replay_buffer(&self) -> &Arc<Mutex<ReplayBuffer>> {
|
|
&self.replay_buffer
|
|
}
|
|
|
|
/// Get the configuration.
|
|
pub const fn config(&self) -> &RainbowAgentConfig {
|
|
&self.config
|
|
}
|
|
}
|
|
|
|
// Manual Debug implementation
|
|
impl std::fmt::Debug for RainbowAgent {
|
|
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
|
let metrics = self.metrics.read().map_err(|_e| std::fmt::Error)?;
|
|
let step_count = *self.step_count.lock().map_err(|_e| std::fmt::Error)?;
|
|
|
|
f.debug_struct("RainbowAgent")
|
|
.field("step_count", &step_count)
|
|
.field("replay_buffer_size", &metrics.replay_buffer_size)
|
|
.field("total_steps", &metrics.total_steps)
|
|
.field("current_loss", &metrics.current_loss)
|
|
.finish_non_exhaustive()
|
|
}
|
|
}
|