Files
foxhunt/docs/plans/2026-03-03-rl-agent-gaps-implementation.md
jgrusewski 866da8bbbb docs: RL agent gaps implementation plan — 7 tasks, vertical slices
Position-aware state, regime detection, slippage, curriculum learning,
multi-timeframe fusion, online learning with EWC, A/B testing extension.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-03-03 02:30:00 +01:00

17 KiB
Raw Blame History

RL Agent Gaps Implementation Plan

For Claude: REQUIRED SUB-SKILL: Use superpowers:executing-plans to implement this plan task-by-task.

Goal: Implement 7 features to close training quality and production readiness gaps in DQN/PPO agents.

Architecture: Vertical slices in dependency order. Each task is a self-contained feature that builds on previous ones. State vector expands from 54 → 57 → 185 dims across tasks.

Tech Stack: Rust, Candle (ML), tokio (async), Prometheus (metrics), gRPC (service integration). Build: SQLX_OFFLINE=true cargo check --workspace. Test: SQLX_OFFLINE=true cargo test -p ml --lib.

Key Codebase Facts:

  • State vector: 51 market features (features/extraction.rs:51 FeatureVector) + 3 portfolio ([value, position_size, spread]) = 54 dims
  • Portfolio features already include position_size (signed). Missing: unrealized_pnl, bars_in_position, cost_basis.
  • RegimeType::classify_from_features() in regime_conditional.rs:101-139 already works (ADX+CUSUM). Only RegimeDetectionEngine::detect_regime() is a stub.
  • deployment/ab_testing.rs already exists (26KB, 730+ lines) with TrafficSplitter. Needs Thompson Sampling extension, not rebuild.
  • Feature dim constant: walk_forward.rs has FEATURE_DIM = 51 (market only, separate from portfolio).
  • Clippy: #![deny(clippy::unwrap_used, clippy::expect_used, clippy::panic, clippy::indexing_slicing)]
  • All struct literals must use ..Default::default() for forward compat.

Task 1: Position-Aware State Encoding

Files:

  • Create: crates/ml/src/features/position_features.rs
  • Modify: crates/ml/src/features/mod.rs — add pub mod position_features;
  • Modify: crates/ml/src/dqn/reward.rs — update TradingState to carry new features
  • Modify: crates/ml/src/walk_forward.rs — update NormStats if needed
  • Modify: crates/ml/src/trainers/dqn/trainer.rs — wire position features into state construction
  • Modify: crates/ml/src/trainers/ppo.rs — same for PPO trajectory state
  • Test: inline #[cfg(test)] in position_features.rs

Context: The state vector is currently 54 dims (51 market + 3 portfolio: [value, position_size, spread]). We add 3 more portfolio features: unrealized_pnl (normalized by EMA σ), bars_in_position (log-scaled), cost_basis (bps relative to price). New total: 57 dims.

Step 1: Write the failing test

Create crates/ml/src/features/position_features.rs with test:

#[cfg(test)]
mod tests {
    use super::*;

    #[test]
    fn test_position_features_flat() {
        let pf = PositionFeatures::new();
        let features = pf.extract(0.0, 0.0, 0, 0.0, 100.0);
        assert_eq!(features.len(), 3);
        assert!((features[0]).abs() < 1e-8); // unrealized_pnl = 0 when flat
        assert!((features[1]).abs() < 1e-8); // bars_in_position = 0
        assert!((features[2]).abs() < 1e-8); // cost_basis = 0 when flat
    }

    #[test]
    fn test_position_features_long() {
        let mut pf = PositionFeatures::new();
        // Feed some volatility history for normalization
        for _ in 0..100 { pf.update_volatility(0.02); }
        let features = pf.extract(1.0, 500.0, 50, 99.0, 100.0);
        assert_eq!(features.len(), 3);
        assert!(features[0] > 0.0); // positive unrealized_pnl
        assert!(features[1] > 0.0); // bars_in_position > 0
        assert!(features[2] > 0.0); // cost_basis: price rose from 99 to 100
    }
}

Step 2: Run test to verify it fails

Run: SQLX_OFFLINE=true cargo test -p ml --lib -- features::position_features Expected: FAIL (module doesn't exist yet or struct not defined)

Step 3: Write minimal implementation

//! Position-aware features for RL agent state encoding.
//! Adds unrealized_pnl, bars_in_position, and cost_basis to the state vector.

/// Extracts position-aware features for the RL state vector.
#[derive(Debug)]
pub struct PositionFeatures {
    vol_ema: f64,
    vol_alpha: f64,
}

impl PositionFeatures {
    pub fn new() -> Self {
        Self { vol_ema: 0.02, vol_alpha: 0.01 }
    }

    pub fn update_volatility(&mut self, abs_return: f64) {
        self.vol_ema = self.vol_alpha * abs_return + (1.0 - self.vol_alpha) * self.vol_ema;
    }

    /// Extract 3 position features:
    /// [0] unrealized_pnl: normalized by EMA volatility
    /// [1] bars_in_position: log(1 + bars) / log(1001), in [0, 1]
    /// [2] cost_basis: (current_price - entry_price) / entry_price * 10000 (bps), clamped
    pub fn extract(
        &self,
        position_size: f64,
        unrealized_pnl: f64,
        bars_held: u64,
        entry_price: f64,
        current_price: f64,
    ) -> Vec<f64> {
        let vol = self.vol_ema.max(1e-8);
        let norm_pnl = if position_size.abs() < 1e-8 { 0.0 } else { unrealized_pnl / vol };
        let norm_bars = (1.0 + bars_held as f64).ln() / (1001.0_f64).ln();
        let cost_bps = if position_size.abs() < 1e-8 || entry_price.abs() < 1e-8 {
            0.0
        } else {
            ((current_price - entry_price) / entry_price * 10000.0).clamp(-500.0, 500.0)
        };
        vec![norm_pnl, norm_bars, cost_bps]
    }
}

impl Default for PositionFeatures {
    fn default() -> Self { Self::new() }
}

Step 4: Run test to verify it passes

Run: SQLX_OFFLINE=true cargo test -p ml --lib -- features::position_features Expected: PASS

Step 5: Wire into DQN trainer

In crates/ml/src/trainers/dqn/trainer.rs, add position_features: PositionFeatures field to DQNTrainer. In the state construction path (where portfolio_features is built), append the 3 new features. Update DQNConfig.state_dim default from 54 to 57 where it's hardcoded.

Step 6: Wire into PPO trainer

Same integration in crates/ml/src/trainers/ppo.rs trajectory collection.

Step 7: Update feature dim references

Grep for 54 and state_dim across the crate. Update hyperopt bounds, walk-forward NormStats, and any hardcoded dims. Use ..Default::default() in all config struct literals.

Step 8: Run full test suite

Run: SQLX_OFFLINE=true cargo test -p ml --lib Expected: All pass (2503+)

Step 9: Clippy + commit

Run: SQLX_OFFLINE=true cargo clippy -p ml -- -D warnings Commit: feat(ml): add position-aware state encoding (54→57 dims)


Task 2: Regime Detection (Feature-Based + HMM)

Files:

  • Modify: crates/ml/src/regime_detection.rs — replace stub with real feature-based classifier
  • Create: crates/ml/src/regime_detection/ directory (split module)
    • mod.rs — re-exports
    • feature_classifier.rs — real-time classifier
    • hmm.rs — offline HMM
  • Modify: crates/ml/src/dqn/regime_conditional.rs — wire RegimeDetectionEngine to use real classifier
  • Test: inline #[cfg(test)] in each new file

Context: RegimeType::classify_from_features() at regime_conditional.rs:101-139 already classifies using ADX (feature 211) and CUSUM (feature 203). The RegimeDetectionEngine wrapping it is the stub. We need to:

  1. Move classification logic into feature_classifier.rs with configurable thresholds
  2. Add HMM for offline validation
  3. Wire the engine to use the real classifier

Step 1: Write failing test for feature classifier

#[test]
fn test_trending_regime_high_adx() {
    let classifier = FeatureClassifier::new(ClassifierConfig::default());
    // ADX > 25, Hurst > 0.55
    let regime = classifier.classify(30.0, 0.6, 0.02, 0.02);
    assert_eq!(regime, RegimeType::Trending);
}

#[test]
fn test_volatile_regime_high_vol() {
    let classifier = FeatureClassifier::new(ClassifierConfig::default());
    // vol > 2σ above mean
    let regime = classifier.classify(10.0, 0.5, 0.08, 0.02);
    assert_eq!(regime, RegimeType::Volatile);
}

Step 2-4: Implement FeatureClassifier

Implement FeatureClassifier with configurable thresholds (ADX trending threshold, Hurst threshold, vol z-score threshold). The classify() method takes computed indicators and returns RegimeType.

Step 5: Implement HMM (3-state)

Simple Baum-Welch HMM with 3 states. fit() takes daily returns, predict() returns most likely state. Store transition matrix for persistence estimation.

Step 6: Wire into RegimeDetectionEngine

Replace stub detect_regime() with call to FeatureClassifier::classify(). Add fit_hmm() and hmm_validate() methods.

Step 7: Wire into RegimeConditionalDQN

In regime_conditional.rs, use the engine instead of inline classify_from_features().

Step 8: Tests + clippy + commit

Run: SQLX_OFFLINE=true cargo test -p ml --lib -- regime Commit: feat(ml): real regime detection with feature classifier + HMM


Task 3: Slippage Modeling

Files:

  • Create: crates/ml/src/risk/slippage.rs
  • Modify: crates/ml/src/risk/mod.rs — add pub mod slippage;
  • Modify: crates/ml/src/dqn/reward.rs — integrate slippage into cost penalty
  • Test: inline #[cfg(test)] in slippage.rs

Context: Current costs are fixed per FactoredAction::transaction_cost() (15/5/10 bps). We add a SlippageModel trait with LinearImpactModel (volume-dependent) and FixedCostModel (backward compat wrapper).

Step 1: Write failing test

#[test]
fn test_linear_impact_higher_volume() {
    let model = LinearImpactModel::new(LinearImpactConfig::default());
    let cost_low = model.estimate_cost(1.0, 1000.0, 100_000.0, OrderType::Market);
    let cost_high = model.estimate_cost(1.0, 10_000.0, 100_000.0, OrderType::Market);
    assert!(cost_high > cost_low, "Higher volume should cost more");
}

Step 2-4: Implement SlippageModel trait + LinearImpactModel + FixedCostModel

pub trait SlippageModel: Send + Sync {
    fn estimate_cost(&self, position_change: f64, volume: f64, adv: f64, order_type: OrderType) -> f64;
}

LinearImpactModel: cost = spread/2 + sqrt_impact * sqrt(volume / adv) + order type premium.

Step 5: Wire into reward.rs

In the transaction cost penalty section (reward.rs:818-876), replace action.transaction_cost() with slippage_model.estimate_cost() when a model is configured. Add slippage_model: Option<Box<dyn SlippageModel>> to RewardFunction.

Step 6: Tests + clippy + commit

Commit: feat(ml): add volume-dependent slippage model


Task 4: Curriculum Learning

Files:

  • Create: crates/ml/src/trainers/curriculum.rs
  • Modify: crates/ml/src/trainers/mod.rs — add pub mod curriculum;
  • Modify: crates/ml/src/trainers/dqn/trainer.rs — integrate curriculum phase into training loop
  • Modify: crates/ml/src/dqn/action_space.rs — add get_curriculum_action_mask(phase, max_position)
  • Test: inline #[cfg(test)]

Context: Uses existing get_valid_action_mask() in action_space.rs:32-58 as foundation. Adds phase-based filtering that restricts action space and regime exposure during early training.

Step 1: Write failing test

#[test]
fn test_phase1_masks_extreme_exposure() {
    let mask = get_curriculum_action_mask(CurriculumPhase::Basic, 2.0);
    // Phase 1: only Flat, Long50, Short50 → exposure indices 1,2,3 valid (not 0=Short100, 4=Long100)
    // Short100 actions (indices 0-8) should be masked
    assert!(!mask[0]); // Short100/Market/Patient
    // Flat actions (indices 18-26) should be valid
    assert!(mask[18]); // Flat/Market/Patient
}

Step 2-4: Implement CurriculumScheduler

pub enum CurriculumPhase { Basic, FullPosition, AllRegimes }

pub struct CurriculumScheduler {
    current_phase: CurriculumPhase,
    sharpe_history: VecDeque<f64>,
    config: CurriculumConfig,
}

Phase transitions gated by Sharpe thresholds over consecutive folds. Phase tracked in checkpoint metadata.

Step 5: Wire into DQN trainer

At epoch boundary, check phase transitions. Apply curriculum mask in addition to position mask.

Step 6: Tests + clippy + commit

Commit: feat(ml): add curriculum learning with 3-phase schedule


Task 5: Multi-Timeframe Feature Fusion (Default On)

Files:

  • Create: crates/ml/src/features/multi_timeframe.rs
  • Create: crates/ml/src/features/bar_resampler.rs
  • Modify: crates/ml/src/features/mod.rs — add modules
  • Modify: crates/ml/src/trainers/dqn/trainer.rs — fuse multi-timeframe into state
  • Modify: crates/ml/src/trainers/ppo.rs — same
  • Test: inline #[cfg(test)]

Context: Currently 1-minute bars only. Resample on-the-fly to 5m/15m/1h. Per-timeframe LSTM encoder → concat → linear projection → 128-dim macro context. Total state: 57 (position-aware) + 128 (multi-tf) = 185 dims. Default on.

Step 1: Implement BarResampler

Aggregates 1m OHLCV bars into higher timeframes (5m, 15m, 1h) with correct OHLC aggregation.

Step 2: Implement MultiTimeframeEncoder

Candle LSTM encoder per timeframe (shared weights optional). Concat 4×64=256, linear project to 128.

pub struct MultiTimeframeEncoder {
    encoders: [LSTMEncoder; 4], // 1m, 5m, 15m, 1h
    projection: Linear,         // 256 → 128
    resampler: BarResampler,
}

Step 3: Wire into trainers

Append 128-dim macro context after 57-dim micro features. Update state_dim configs.

Step 4: Tests + clippy + commit

Commit: feat(ml): add multi-timeframe feature fusion (default on, 57→185 dims)


Task 6: Online Learning with EWC

Files:

  • Create: crates/ml/src/trainers/online_learning.rs
  • Modify: crates/ml/src/trainers/mod.rs — add module
  • Modify: crates/ml/src/trainers/dqn/config.rs — add OnlineLearningConfig
  • Test: inline #[cfg(test)]

Context: After walk-forward training, the model receives live trade data. Every 100 trades, perform a mini-update with EWC regularization to prevent catastrophic forgetting.

Step 1: Implement EWC (Elastic Weight Consolidation)

pub struct EWCRegularizer {
    fisher_diagonal: HashMap<String, Tensor>,  // per-parameter Fisher info
    optimal_params: HashMap<String, Tensor>,   // θ* from initial training
    lambda: f64,                                // EWC strength (default 1000.0)
}

Compute Fisher Information Matrix (diagonal approx) from initial training data. EWC loss = task_loss + λ * Σ F_i * (θ_i - θ*_i)².

Step 2: Implement OnlineLearner

pub struct OnlineLearner {
    rolling_buffer: VecDeque<Experience>,
    buffer_capacity: usize,        // 10,000
    update_interval: usize,        // 100 trades
    trade_count: usize,
    ewc: EWCRegularizer,
    lr_multiplier: f64,            // 0.1
    max_grad_norm: f64,            // 1.0
    sharpe_monitor: RollingSharpe, // 500-trade window
    checkpoint_path: PathBuf,
}

Step 3: Safety rails

  • Gradient clipping: max_grad_norm = 1.0
  • Learning rate: base_lr * 0.1
  • Auto-rollback: if rolling Sharpe degrades > 20% over 500 trades
  • Kill switch: if rolling Sharpe < -1.0, freeze model

Step 4: Tests + clippy + commit

Commit: feat(ml): add online learning with EWC regularization


Task 7: A/B Testing — Thompson Sampling Extension

Files:

  • Modify: crates/ml/src/deployment/ab_testing.rs — add Thompson Sampling strategy
  • Test: extend existing tests in same file

Context: ab_testing.rs already has 730+ lines with TrafficSplitter, ABTestConfig, TestGroup, and 4 splitting strategies (HashBased, Random, RoundRobin, WeightedRandom). We ADD ThompsonSampling as a 5th strategy with Bayesian promotion criteria.

Step 1: Add ThompsonSampling variant to TrafficSplittingStrategy

pub enum TrafficSplittingStrategy {
    HashBased,
    Random,
    RoundRobin,
    WeightedRandom,
    ThompsonSampling,  // NEW
}

Step 2: Implement Thompson Sampling logic

pub struct ThompsonSamplingState {
    control_alpha: f64,    // Beta distribution α for champion
    control_beta: f64,     // Beta distribution β for champion
    treatment_alpha: f64,  // Beta distribution α for challenger
    treatment_beta: f64,   // Beta distribution β for challenger
}

Sample from Beta(α, β) for each arm; assign trade to higher sample. Update α/β after trade outcome.

Step 3: Add Bayesian promotion criteria

pub struct PromotionCriteria {
    pub min_trades: usize,           // 500
    pub bayesian_confidence: f64,    // 0.95
    pub max_drawdown_ratio: f64,     // 1.2 (challenger DD ≤ 1.2× champion DD)
    pub require_manual_approval: bool, // true
}

Step 4: Wire into TrafficSplitter::assign_group()

Add ThompsonSampling arm to the match in assign_group() (currently at line 443-448).

Step 5: Tests + clippy + commit

Commit: feat(ml): add Thompson Sampling and Bayesian promotion to A/B testing


Execution Notes

Build command: SQLX_OFFLINE=true cargo check --workspace Test command: SQLX_OFFLINE=true cargo test -p ml --lib Clippy command: SQLX_OFFLINE=true cargo clippy -p ml -- -D warnings

Critical rules:

  • Never use .unwrap(), .expect(), .panic!(), or [index] — use .get(), ?, .ok_or()
  • All struct literals use ..Default::default()
  • Candle f32 trap: Tensor * 0.5 fails. Use tensor.broadcast_mul(&Tensor::new(0.5_f32, device)?)
  • Each task should be committed individually before starting the next
  • Run full cargo test -p ml --lib after each task to catch regressions