Files
foxhunt/docs/codebase-cleanup/PER_DATA_FLOW_DIAGRAM.txt
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

236 lines
13 KiB
Plaintext
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
================================================================================
PRIORITIZED EXPERIENCE REPLAY - DATA FLOW DIAGRAM
================================================================================
┌─────────────────────────────────────────────────────────────────────────────┐
│ DQN TRAINING LOOP │
└─────────────────────────────────────────────────────────────────────────────┘
┌────────────────────┐
│ 1. Store │
│ Experience │
│ (state, action, │
│ reward, next, │
│ done) │
└──────┬─────────────┘
v
┌────────────────────────────────────────────────────────┐
│ PrioritizedReplayBuffer::push() │
│ • Initial priority = max_priority (ensures sampling) │
│ • Segment tree update: O(log n) │
└──────┬─────────────────────────────────────────────────┘
v
┌────────────────────┐
│ 2. Sample Batch │
│ (batch_size=32) │
└──────┬─────────────┘
v
┌──────────────────────────────────────────────────────────────┐
│ PrioritizedReplayBuffer::sample() │
│ ┌────────────────────────────────────────────────────────┐ │
│ │ For each sample i: │ │
│ │ │ │
│ │ 1. Get priority: p_i from segment tree │ │
│ │ │ │
│ │ 2. Calculate probability: │ │
│ │ P(i) = p_i / Σ(p_j) │ │
│ │ │ │
│ │ 3. Calculate current beta (with annealing): │ │
│ │ progress = step / beta_annealing_steps │ │
│ │ β = β_start + (β_max - β_start) × progress │ │
│ │ β ∈ [0.4, 1.0] │ │
│ │ │ │
│ │ 4. Calculate IS weight: │ │
│ │ w_raw = (N × P(i))^(-β) │ │
│ │ w_i = w_raw / max(w_j) ← normalization │ │
│ └────────────────────────────────────────────────────────┘ │
│ │
│ Returns: (experiences, weights, indices) │
└──────┬───────────────────────────────────────────────────────┘
v
┌────────────────────┐
│ 3. Forward Pass │
│ • Main network │
│ • Target network │
└──────┬─────────────┘
v
┌──────────────────────────────────────────────────────────────┐
│ WorkingDQN::train_step() │
│ ┌────────────────────────────────────────────────────────┐ │
│ │ 1. Compute TD errors: │ │
│ │ δ_i = Q(s,a) - (r + γ max Q'(s',a')) │ │
│ │ │ │
│ │ 2. Apply IS weights to loss: │ │
│ │ │ │
│ │ Standard Loss (MSE/Huber): │ │
│ │ weighted_diff = (Q - target) × w_i │ │
│ │ loss = mean(weighted_diff²) │ │
│ │ │ │
│ │ Distributional Loss (C51): │ │
│ │ per_sample_loss = -Σ target_i × log(pred_i) │ │
│ │ weighted_loss = per_sample_loss × w_i │ │
│ │ loss = mean(weighted_loss) │ │
│ └────────────────────────────────────────────────────────┘ │
└──────┬───────────────────────────────────────────────────────┘
v
┌────────────────────┐
│ 4. Backward Pass │
│ • Gradients │
│ • Optimizer step │
└──────┬─────────────┘
v
┌──────────────────────────────────────────────────────────────┐
│ 5. Update Priorities │
│ ┌────────────────────────────────────────────────────────┐ │
│ │ For each sampled index i: │ │
│ │ │ │
│ │ 1. Get TD error: δ_i │ │
│ │ │ │
│ │ 2. Calculate new priority: │ │
│ │ p_i = |δ_i| + ε (ε = 1e-6) │ │
│ │ │ │
│ │ 3. Update segment tree: O(log n) │ │
│ │ memory.update_priorities(indices, priorities) │ │
│ └────────────────────────────────────────────────────────┘ │
└──────┬───────────────────────────────────────────────────────┘
v
┌────────────────────┐
│ 6. Step Beta │
│ memory.step() │
│ training_step++ │
└────────────────────┘
================================================================================
SEGMENT TREE STRUCTURE
================================================================================
Example for capacity = 8:
Tree array indices:
[0] [1] [2] [3] [4] [5] [6] [7] [8] [9] [10][11][12][13][14][15]
x ROOT L0 R0 L1 R1 L2 R2 EXP EXP EXP EXP EXP EXP EXP EXP
↑ ↑ ↑ ↑ ↑ ↑ ↑ ↑ ↑
│ └───┴───┴───┴───┴───┴───┴───┘
│ Leaf nodes (experiences 0-7)
Sum of all priorities
Internal nodes store cumulative sums:
- tree[1] = total priority sum
- tree[2] = sum of left subtree (experiences 0-3)
- tree[3] = sum of right subtree (experiences 4-7)
- tree[i] = tree[2i] + tree[2i+1]
Sampling algorithm (O(log n)):
1. Generate random value v ∈ [0, tree[1]]
2. Start at root (idx=1)
3. While not at leaf:
- If v <= tree[left_child]: go left
- Else: subtract tree[left_child] from v, go right
4. Return (idx - capacity) as experience index
Update algorithm (O(log n)):
1. Set tree[idx + capacity] = new_priority
2. Propagate up: tree[parent] = tree[left] + tree[right]
3. Repeat until root
================================================================================
BETA ANNEALING SCHEDULE
================================================================================
Training Step │ Progress │ Beta │ IS Correction Strength
──────────────┼──────────┼──────────┼────────────────────────
0 │ 0.0% │ 0.40 │ Minimal (exploration)
50,000 │ 10.0% │ 0.46 │ Growing
100,000 │ 20.0% │ 0.52 │ ↓
150,000 │ 30.0% │ 0.58 │ ↓
200,000 │ 40.0% │ 0.64 │ ↓
250,000 │ 50.0% │ 0.70 │ ↓
300,000 │ 60.0% │ 0.76 │ ↓
350,000 │ 70.0% │ 0.82 │ ↓
400,000 │ 80.0% │ 0.88 │ ↓
450,000 │ 90.0% │ 0.94 │ ↓
500,000 │ 100.0% │ 1.00 │ Full correction (convergence)
Formula: β = β_start + (β_max - β_start) × min(1.0, step / annealing_steps)
================================================================================
IMPORTANCE SAMPLING WEIGHT CALCULATION
================================================================================
Given:
- N = buffer size (e.g., 100,000)
- P(i) = sampling probability for experience i
- β = current beta value (0.4 → 1.0)
Step-by-step calculation:
1. Raw weight:
w_raw(i) = (N × P(i))^(-β)
2. Find maximum weight:
P_min = min(P(j)) for all j
w_max = (N × P_min)^(-β)
3. Normalize weight:
w_i = w_raw(i) / w_max
4. Clamp to reasonable range:
w_i = min(w_i, 10.0) ← prevents extreme values
Properties:
- w_i ∈ [0, 1] after normalization
- Higher priority → higher P(i) → lower w_i (compensates for bias)
- Lower priority → lower P(i) → higher w_i (upweights rare samples)
- β=0 → w_i=1 (no correction, pure prioritization)
- β=1 → full correction (unbiased gradient estimates)
================================================================================
PRIORITY UPDATE EXAMPLES
================================================================================
Example 1: High TD error (important transition)
TD error: δ = 5.0
Priority: p = |5.0| + 1e-6 = 5.000001
Result: High sampling probability in next batch
Example 2: Low TD error (well-learned transition)
TD error: δ = 0.01
Priority: p = |0.01| + 1e-6 = 0.010001
Result: Low sampling probability
Example 3: New experience (no TD error yet)
TD error: N/A
Priority: p = max_priority (e.g., 10.0)
Result: Guaranteed to be sampled at least once
================================================================================
MEMORY EFFICIENCY
================================================================================
For 1M capacity buffer:
Component Memory Usage
─────────────────────────────────────────
Segment tree 8 MB (2 × 1M × 4 bytes)
Experience buffer Varies (depends on state size)
Atomic counters 32 bytes
RNG state ~100 bytes
─────────────────────────────────────────
Total PER overhead ~8 MB + experiences
Comparison to uniform buffer:
- Uniform: experiences only
- PER: experiences + 8 MB overhead
- Overhead: ~0.8% for typical state sizes
================================================================================