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>
5.3 KiB
5.3 KiB
WAVE 26 P1.2: Multi-Head Self-Attention Implementation Report
Summary
Implemented multi-head self-attention layer for temporal pattern recognition in DQN architecture.
Files Changed
Created Files
/home/jgrusewski/Work/foxhunt/ml/src/dqn/attention.rs(580 lines)- Complete multi-head attention implementation
- Scaled dot-product attention with optional masking
- Xavier initialization for all linear layers
- Optional layer normalization and residual connections
- Comprehensive TDD test suite (8 tests)
Modified Files
/home/jgrusewski/Work/foxhunt/ml/src/dqn/mod.rs- Added
pub mod attention;declaration (line 10) - Added
pub use attention::{MultiHeadAttention, MultiHeadAttentionConfig};(line 62)
- Added
Implementation Details
Architecture
Input (batch, seq_len, embed_dim)
|
├─> Query (WQ) ─┐
├─> Key (WK) ───┤
└─> Value (WV) ─┴─> Scaled Dot-Product Attention
|
v
Multi-Head Concat
|
v
Output Linear (WO)
|
v
Output (batch, seq_len, embed_dim)
Key Features
-
Multi-Head Attention
- Configurable number of heads (default: 4)
- Configurable embedding dimension (default: 64)
- Automatic head dimension calculation:
head_dim = embed_dim / num_heads
-
Scaled Dot-Product Attention
- Formula:
Attention(Q, K, V) = softmax(QK^T / √d_k) V - Scaling prevents gradient saturation for large dimensions
- Optional attention masking for causal/padding masks
- Formula:
-
Initialization & Stability
- Xavier/Glorot initialization for all linear layers
- Layer normalization for training stability (optional)
- Residual connections for gradient flow (optional)
-
Configuration Options
MultiHeadAttentionConfig { embed_dim: 64, // Must be divisible by num_heads num_heads: 4, // Number of attention heads dropout: 0.1, // Dropout probability use_layer_norm: true, // Enable layer normalization layer_norm_eps: 1e-5, // LayerNorm epsilon use_residual: true, // Enable residual connections }
TDD Test Coverage
Created 8 comprehensive tests before implementation:
-
test_config_validation- Validates embed_dim > 0
- Validates num_heads > 0
- Validates embed_dim divisible by num_heads
-
test_default_config- Verifies default configuration values
- Checks head_dim calculation
-
test_attention_creation- Tests successful layer instantiation
- Validates configuration propagation
-
test_forward_pass_shape- Input:
(batch=2, seq_len=8, embed_dim=64) - Output:
(batch=2, seq_len=8, embed_dim=64) - Verifies shape preservation
- Input:
-
test_forward_with_mask- Tests causal mask application (lower triangular)
- Mask format:
0.0for attend,-inffor mask - Verifies masked attention computation
-
test_dimension_mismatch- Tests error handling for wrong input dimensions
- Verifies
MLError::DimensionMismatcherror
-
test_residual_connection- Tests residual connection functionality
- Validates output shape with residuals
-
test_multiple_heads- Tests with 1, 2, 4, 8 heads
- Validates multi-head parallelization
Error Handling
MLError::ConfigurationError: Invalid configuration (divide by zero, etc.)MLError::DimensionMismatch: Input shape mismatchMLError::InitializationError: Failed parameter initializationMLError::ModelError: Forward pass failuresMLError::TensorOperationError: Tensor manipulation failures
Integration Path
The attention layer can be integrated into network architectures as follows:
use ml::dqn::{MultiHeadAttention, MultiHeadAttentionConfig};
use candle_nn::VarBuilder;
// Create configuration
let config = MultiHeadAttentionConfig::new(64, 4)?;
// Initialize attention layer
let attention = MultiHeadAttention::new(config, &var_builder, &device)?;
// Forward pass (no mask)
let output = attention.forward(&input, None)?;
// Forward pass with causal mask
let causal_mask = create_causal_mask(seq_len, &device)?;
let output = attention.forward(&input, Some(&causal_mask))?;
Performance Characteristics
- Memory: O(batch_size × seq_len² × num_heads) for attention scores
- Computation: O(batch_size × seq_len² × embed_dim × num_heads)
- GPU Acceleration: Full CUDA support via candle_core
- Numerical Stability: Xavier initialization + optional LayerNorm
Next Steps
-
Integration Testing
- Integrate into QNetwork architecture
- Test with DQN training loop
- Validate gradient flow through attention
-
Performance Optimization
- Profile attention computation
- Benchmark vs baseline DQN
- Optimize for different sequence lengths
-
Hyperparameter Tuning
- Optimal number of heads for trading
- Optimal embedding dimension
- Dropout rate tuning
References
- Vaswani et al., "Attention Is All You Need" (2017)
- Xavier Glorot initialization for gradient stability
- Layer normalization for training dynamics (Ba et al., 2016)