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>
5.3 KiB
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
/// 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
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
max_validation_batches: opts.max_validation_batches,
2. /home/jgrusewski/Work/foxhunt/ml/src/trainers/tft.rs
Line 449-452: Added field to TFTTrainerConfig
/// 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
max_validation_batches: None, // Default: unlimited (use all validation data)
Line 522-523: Added to to_training_config() method
validation_batch_size: self.validation_batch_size,
max_validation_batches: self.max_validation_batches,
Line 1468-1470: Modified validation loop to limit batches
// 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
// 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
/// 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
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
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
max_validation_batches: None, // Default: unlimited validation
Usage
Command Line
# 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
# 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
-
Two Config Structs: The implementation spans both
TFTTrainerConfig(high-level API) andTFTTrainingConfig(internal training state). -
Default Behavior: When
max_validation_batchesisNone, the system processes all validation batches (backward compatible). -
Memory Savings: Each validation batch uses ~10MB, so limiting to 50 batches saves ~1260MB (126 batches × 10MB).
-
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:
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.