Files
foxhunt/docs/dqn_refactoring_implementation.md
jgrusewski 2df1ea92e1 feat(ml): WAVE 29 DQN Codebase Cleanup & Refactoring Campaign
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>
2025-11-27 23:46:13 +01:00

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.rs with 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:

  1. Automated extraction scripts (provided above)
  2. Incremental verification with cargo check after each module
  3. Careful import management (40+ imports needed for trainer_core)
  4. 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.