BREAKING CHANGES: - Removed orphaned dqn.rs monolithic trainer (4,975 lines) - Removed orphaned dqn_ensemble.rs module (816 lines) - Removed orphaned tft.rs and tft_complete_int8_integration_test.rs - TFT trainer split into modular directory structure DQN Module Refactoring: - Split trainers/dqn.rs into modular structure (config.rs, statistics.rs, trainer.rs) - Fixed hyperopt 39D search space (continuous params only) - Boolean flags (use_dueling, use_double_dqn, use_per, use_noisy_nets) are now FIXED architectural decisions - use_distributional defaults to false (Candle BUG #36 - scatter_add gradient issues) Clean Module Structure: - ml/src/trainers/dqn/ directory with proper mod.rs exports - ml/src/trainers/tft/ directory with config.rs, types.rs, model.rs, trainer.rs, tests.rs - All P0 features validated: TD-error clamping, batch diversity, LR scheduler, priority staleness Documentation: - Added comprehensive docs in docs/codebase-cleanup/ - ADR-001 for DQN refactoring decisions - Rainbow DQN component matrix and quick reference guides Build Status: Compiles with zero errors 🤖 Generated with [Claude Code](https://claude.com/claude-code) Co-Authored-By: Claude <noreply@anthropic.com>
8.9 KiB
8.9 KiB
DQN Refactoring Implementation Guide
Current Status: READY FOR EXECUTION
Files Analysis
| File | Current Lines | Target | Status |
|---|---|---|---|
trainers/dqn.rs |
4,975 | <1,000 each | BACKUP CREATED |
hyperopt/adapters/dqn.rs |
3,162 | <1,000 each | Pending |
trainers/tft.rs |
2,915 | <1,000 each | Pending |
trainers/mamba2.rs |
544 | OK (under 1K) | ✅ No action needed |
Extraction Map for dqn.rs
Module 1: dqn/config.rs (800 lines)
Source Lines: 48-747
Contents:
// Constants and type aliases (lines 48-56)
const EPISODE_LENGTH: usize = 200;
type FeatureVector = [f64; 54];
type FeatureVector51 = [f64; 51];
// FeatureStatistics (lines 77-157)
pub struct FeatureStatistics { /* Welford's algorithm */ }
// DQNHyperparameters (lines 425-610)
pub struct DQNHyperparameters { /* 60+ fields */ }
// DQNHyperparameters impl (lines 615-747)
impl DQNHyperparameters {
pub fn conservative() -> Self { /* ... */ }
}
Required Imports:
// From trainers/mod.rs
use crate::trainers::TargetUpdateMode;
Module 2: dqn/agent_wrapper.rs (260 lines)
Source Lines: 164-424
Contents:
// DQNAgentType enum (lines 164-169)
pub enum DQNAgentType { /* Standard | RegimeConditional */ }
// QValueStats (lines 182-194)
pub struct QValueStats { /* C51 adaptive bounds */ }
// DQNAgentType impl (lines 196-421)
impl DQNAgentType {
// 20+ unified API methods
}
Required Imports:
use candle_core::{Device, Tensor};
use crate::dqn::{
dqn::{WorkingDQN},
regime_conditional::{RegimeConditionalDQN, RegimeMetrics, RegimeType},
action_space::FactoredAction,
Experience,
};
use crate::MLError;
Module 3: dqn/training_monitor.rs (265 lines)
Source Lines: 728-1012
Contents:
// TrainingMonitor struct (lines 728-746)
struct TrainingMonitor { /* validation fields */ }
// TrainingMonitor impl (lines 748-1012)
impl TrainingMonitor {
fn new(epoch: usize) -> Self { /* ... */ }
fn track_reward(&mut self, reward: f32) { /* ... */ }
fn track_action(&mut self, action: &FactoredAction) { /* ... */ }
fn validate_rewards(&mut self) -> Result<()> { /* ... */ }
fn validate_action_diversity(&self) -> Result<()> { /* ... */ }
fn validate_q_value_balance(&self) -> Result<()> { /* ... */ }
fn validate_all(&mut self) -> Result<()> { /* ... */ }
// ... more validation methods
}
Required Imports:
use anyhow::Result;
use tracing::{info, warn};
use crate::dqn::action_space::FactoredAction;
Module 4: dqn/trainer_core.rs (600 lines estimated)
Source Lines: 1013-1680 (approx)
Contents:
// DQNTrainer struct (lines 1013-1116)
pub struct DQNTrainer {
agent: Arc<RwLock<DQNAgentType>>,
hyperparams: DQNHyperparameters,
device: Device,
// ... 20+ fields
}
// Debug impl (lines 1118-1124)
impl std::fmt::Debug for DQNTrainer { /* ... */ }
// Constructor impl (lines 1126-1450)
impl DQNTrainer {
pub fn new(hyperparams: DQNHyperparameters) -> Result<Self> { /* ... */ }
pub fn new_with_debug(hyperparams: DQNHyperparameters, debug_logging: bool) -> Result<Self> { /* ... */ }
pub fn with_feature_cache(mut self, cache_dir: PathBuf) -> Self { /* ... */ }
}
Required Imports: (See full file - 40+ imports needed)
Module 5: dqn/training_loop.rs (1500 lines estimated)
Source Lines: 1464-2980 (main train methods)
Contents:
impl DQNTrainer {
pub async fn train<F>(&mut self, ...) -> Result<TrainingMetrics> {
// Main training loop: lines 1464-2800
}
pub async fn train_from_parquet<F>(&mut self, ...) -> Result<TrainingMetrics> {
// Parquet training loop: lines 2981-3400
}
// Helper methods:
fn calculate_epoch_metrics(...) -> Result<...> { /* ... */ }
fn calculate_adaptive_bounds(...) -> (f64, f64) { /* ... */ }
fn check_early_stopping(...) -> Option<String> { /* ... */ }
}
Module 6: dqn/data_loading.rs (600 lines estimated)
Source Lines: 3482-4100 (OHLCV extraction, feature creation)
Contents:
impl DQNTrainer {
pub fn extract_ohlcv_bars_from_dbn(&self, file_path: &Path) -> Result<Vec<OHLCVBar>> {
// DBN file parsing: lines 3482-3600
}
fn create_features(&self, ...) -> Result<Vec<f32>> {
// Feature engineering: lines 3609-4100
}
}
Module 7: dqn/checkpointing.rs (200 lines estimated)
Source Lines: Scattered throughout, needs extraction
Contents:
impl DQNTrainer {
fn save_checkpoint(&self, epoch: usize) -> Result<()> { /* ... */ }
fn load_checkpoint(&mut self, path: &str) -> Result<()> { /* ... */ }
}
Module 8: dqn/mod.rs (50 lines)
New file
Contents:
//! DQN Trainer Module
//!
//! Modularized DQN trainer for maintainability and testability.
//! This module preserves the original public API via re-exports.
// Internal modules
mod config;
mod agent_wrapper;
mod training_monitor;
mod trainer_core;
mod training_loop;
mod data_loading;
mod checkpointing;
// Public re-exports (maintain API compatibility)
pub use config::{
DQNHyperparameters,
FeatureStatistics,
FeatureVector,
FeatureVector51,
EPISODE_LENGTH,
};
pub use agent_wrapper::{DQNAgentType, QValueStats};
pub use trainer_core::DQNTrainer;
// TrainingMonitor is internal, not part of public API
pub(crate) use training_monitor::TrainingMonitor;
Implementation Scripts
Script 1: Create dqn/config.rs
#!/bin/bash
# Extract config module from dqn.rs
SOURCE="ml/src/trainers/dqn.rs.backup"
TARGET="ml/src/trainers/dqn/config.rs"
cat > "$TARGET" <<'HEADER'
//! DQN Configuration and Hyperparameters
//!
//! Contains all configuration types for DQN training including:
//! - Constants and type aliases
//! - Feature normalization statistics (Welford's algorithm)
//! - DQN hyperparameters (60+ tunable parameters)
use crate::trainers::TargetUpdateMode;
HEADER
# Extract lines 48-747 (constants, FeatureStatistics, DQNHyperparameters)
sed -n '48,747p' "$SOURCE" >> "$TARGET"
echo "✅ Created $TARGET"
Script 2: Create dqn/agent_wrapper.rs
#!/bin/bash
# Extract agent wrapper module from dqn.rs
SOURCE="ml/src/trainers/dqn.rs.backup"
TARGET="ml/src/trainers/dqn/agent_wrapper.rs"
cat > "$TARGET" <<'HEADER'
//! DQN Agent Type Wrapper
//!
//! Provides unified API for both standard and regime-conditional DQN agents.
//! Allows transparent switching between single-head and multi-head Q-networks.
use candle_core::{Device, Tensor};
use candle_nn::VarMap;
use crate::dqn::action_space::FactoredAction;
use crate::dqn::dqn::WorkingDQN;
use crate::dqn::regime_conditional::{RegimeConditionalDQN, RegimeMetrics, RegimeType};
use crate::dqn::Experience;
use crate::MLError;
HEADER
# Extract lines 164-424 (DQNAgentType, QValueStats)
sed -n '164,424p' "$SOURCE" >> "$TARGET"
echo "✅ Created $TARGET"
Script 3: Create dqn/training_monitor.rs
#!/bin/bash
# Extract training monitor module from dqn.rs
SOURCE="ml/src/trainers/dqn.rs.backup"
TARGET="ml/src/trainers/dqn/training_monitor.rs"
cat > "$TARGET" <<'HEADER'
//! Training Monitoring and Validation
//!
//! Validates training progress to prevent common bugs:
//! - Constant rewards (reward shaping issues)
//! - Action collapse (policy degeneration)
//! - Q-value imbalance (exploration failure)
use anyhow::{Context, Result};
use tracing::{info, warn};
use crate::dqn::action_space::FactoredAction;
HEADER
# Extract lines 728-1012 (TrainingMonitor)
sed -n '728,1012p' "$SOURCE" >> "$TARGET"
echo "✅ Created $TARGET"
Execution Checklist
- Backup created:
dqn.rs.backup - Directories created:
dqn/,tft/,mamba2/,shared/ - ADR written:
docs/ADR-001-dqn-refactoring.md - Extract
dqn/config.rs - Extract
dqn/agent_wrapper.rs - Extract
dqn/training_monitor.rs - Verify:
cargo check --package ml - Extract
dqn/trainer_core.rs - Extract
dqn/training_loop.rs - Extract
dqn/data_loading.rs - Extract
dqn/checkpointing.rs - Create
dqn/mod.rswith re-exports - Delete original
dqn.rs - Verify:
cargo check --package ml - Run tests:
cargo test --package ml --lib - Verify all 19 DQN tests pass
- Verify hyperopt adapter compiles
Notes for Next Context
Due to file size (4,975 lines) and context limits, this refactoring requires:
- Automated extraction scripts (provided above)
- Incremental verification with
cargo checkafter each module - Careful import management (40+ imports needed for trainer_core)
- Test validation after completion
The architecture is sound. Implementation is mechanical but requires attention to:
- Import statements (many cross-dependencies)
- Method visibility (pub vs pub(crate))
- Re-exports maintaining public API
Next agent should execute extraction scripts in order, verifying build after each step.