Files
foxhunt/crates/ml-dqn/src/rainbow_agent.rs
jgrusewski 49602a93a6 fix(clippy): ml-dqn — 148 errors fixed
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>
2026-03-19 00:49:38 +01:00

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()
}
}