Files
foxhunt/MAX_VALIDATION_BATCHES_IMPLEMENTATION.md
jgrusewski 0fced7619d fix(ml): Add max_validation_batches to prevent validation OOM
Workaround for Candle's lack of CUDA memory clearing APIs

Problem:
- Validation needs 1760MB for 176 batches
- Only 2485MB available after training
- Candle doesn't expose cuda::empty_cache() to free optimizer memory
- Result: OOM during validation despite optimizer drop

Solution:
- Add max_validation_batches parameter to limit validation batches
- Default: None (unlimited, backward compatible)
- Recommended for 4GB GPUs: 50 batches (~500MB vs 1760MB)

Changes:
1. CLI parameter: --max-validation-batches <num>
2. TFTTrainerConfig: max_validation_batches field
3. TFTTrainingConfig: max_validation_batches field
4. Validation loop: .take(max_batches) to limit batches
5. Updated: benchmarks, legacy binary for compatibility

Impact:
- 50 batches: 1611MB + 500MB = 2111MB < 2485MB 
- 176 batches: 1611MB + 1760MB = 3371MB > 2485MB 
- Memory savings: 1260MB (72% reduction)
- Trade-off: Validates on subset (28% of data)

Files Modified:
- ml/examples/train_tft_parquet.rs (CLI + config)
- ml/src/trainers/tft.rs (config + validation loop)
- ml/src/tft/training.rs (internal config)
- ml/src/benchmark/tft_benchmark.rs (compatibility)
- ml/src/bin/train_tft.rs (compatibility)

🤖 Generated with [Claude Code](https://claude.com/claude-code)

Co-Authored-By: Claude <noreply@anthropic.com>
2025-10-26 20:58:49 +01:00

170 lines
5.3 KiB
Markdown
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.
# max_validation_batches Parameter Implementation
## Summary
Added `max_validation_batches` parameter to limit validation memory usage as a workaround for Candle's lack of CUDA memory clearing APIs.
**Problem**: Validation needs 1760MB for 176 batches, but only 2485MB available → OOM
**Solution**: Limit validation to 50 batches → reduces memory to 500MB → fits in available memory
## Files Modified
### 1. `/home/jgrusewski/Work/foxhunt/ml/examples/train_tft_parquet.rs`
**Line 162-166**: Added CLI parameter
```rust
/// Maximum validation batches to run (default: unlimited, use 50 for 4GB GPUs)
/// Limits validation to N batches to reduce memory usage. Each batch uses ~10MB,
/// so 50 batches = ~500MB vs 1760MB for full validation (176 batches).
#[arg(long)]
max_validation_batches: Option<usize>,
```
**Line 197-201**: Added logging for the parameter
```rust
if let Some(max_val_batches) = opts.max_validation_batches {
info!(" • Max validation batches: {} (memory optimization)", max_val_batches);
} else {
info!(" • Max validation batches: unlimited");
}
```
**Line 288**: Added to config construction
```rust
max_validation_batches: opts.max_validation_batches,
```
### 2. `/home/jgrusewski/Work/foxhunt/ml/src/trainers/tft.rs`
**Line 449-452**: Added field to `TFTTrainerConfig`
```rust
/// Maximum validation batches to run (None = unlimited)
/// Limits validation to N batches to reduce memory usage on constrained GPUs.
/// Example: 50 batches = ~500MB vs 1760MB for full validation (176 batches)
pub max_validation_batches: Option<usize>,
```
**Line 482**: Added to Default implementation
```rust
max_validation_batches: None, // Default: unlimited (use all validation data)
```
**Line 522-523**: Added to `to_training_config()` method
```rust
validation_batch_size: self.validation_batch_size,
max_validation_batches: self.max_validation_batches,
```
**Line 1468-1470**: Modified validation loop to limit batches
```rust
// Limit validation batches if max_validation_batches is set (memory optimization)
let max_batches = self.training_config.max_validation_batches.unwrap_or(usize::MAX);
for (i, batch) in val_loader.iter().take(max_batches).enumerate() {
```
**Line 1521-1527**: Added logging when validation is limited
```rust
// Log if validation was limited for memory optimization
if let Some(max) = self.training_config.max_validation_batches {
info!(
"[VALIDATION] Processed {} batches (limited to {} for memory optimization)",
batch_count, max
);
}
```
### 3. `/home/jgrusewski/Work/foxhunt/ml/src/tft/training.rs`
**Line 58-61**: Added field to `TFTTrainingConfig`
```rust
/// Maximum validation batches to run (None = unlimited)
/// Limits validation to N batches to reduce memory usage on constrained GPUs.
/// Example: 50 batches = ~500MB vs 1760MB for full validation (176 batches)
pub max_validation_batches: Option<usize>,
```
**Line 105**: Added to Default implementation
```rust
max_validation_batches: None, // Default: unlimited (use all validation data)
```
### 4. `/home/jgrusewski/Work/foxhunt/ml/src/benchmark/tft_benchmark.rs`
**Line 554**: Added to benchmark config
```rust
max_validation_batches: None, // Benchmark uses all validation data
```
### 5. `/home/jgrusewski/Work/foxhunt/ml/src/bin/train_tft.rs`
**Line 203**: Added to legacy train_tft config
```rust
max_validation_batches: None, // Default: unlimited validation
```
## Usage
### Command Line
```bash
# Train with limited validation (50 batches for 4GB GPUs)
cargo run -p ml --example train_tft_parquet --release --features cuda -- \
--parquet-file test_data/ES_FUT_180d.parquet \
--epochs 50 \
--max-validation-batches 50
# Train with unlimited validation (default)
cargo run -p ml --example train_tft_parquet --release --features cuda -- \
--parquet-file test_data/ES_FUT_180d.parquet \
--epochs 50
```
### Expected Impact
With `--max-validation-batches 50`:
- **Validation memory**: ~500MB (vs 1760MB for 176 batches)
- **Available memory**: 2485MB
- **Total usage**: 1611MB (training) + 500MB (validation) = 2111MB < 2485MB ✅
- **Trade-off**: Validation on subset (28% of data), but training still uses all data
### Verification
```bash
# Check compilation
cargo check -p ml --example train_tft_parquet
# Test the parameter
cargo run -p ml --example train_tft_parquet --release --features cuda -- \
--parquet-file test_data/ES_FUT_small.parquet \
--epochs 3 \
--max-validation-batches 10
```
Expected log output:
```
• Max validation batches: 10 (memory optimization)
...
[VALIDATION] Processed 10 batches (limited to 10 for memory optimization)
```
## Implementation Notes
1. **Two Config Structs**: The implementation spans both `TFTTrainerConfig` (high-level API) and `TFTTrainingConfig` (internal training state).
2. **Default Behavior**: When `max_validation_batches` is `None`, the system processes all validation batches (backward compatible).
3. **Memory Savings**: Each validation batch uses ~10MB, so limiting to 50 batches saves ~1260MB (126 batches × 10MB).
4. **Validation Quality**: With 50 batches, you still validate on ~28% of data, which provides reasonable accuracy estimates while avoiding OOM.
## Testing
All changes compile successfully:
```bash
cargo check -p ml
# Output: Finished `dev` profile [unoptimized + debuginfo] target(s) in 1.27s
```
## Status
**COMPLETE** - All 5 files modified, all compilation errors resolved.