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>
14 KiB
TFT Trainer Split Analysis
Agent: tft-splitter
Date: 2025-11-27
Target File: ml/src/trainers/tft.rs (2,915 lines)
Reference: DQN split (4 files: config.rs, statistics.rs, trainer.rs, mod.rs)
Executive Summary
The TFT trainer module can be split into 6 files following a similar pattern to DQN, but with additional files for QAT metrics and the TFTModel trait abstraction.
Proposed Structure:
ml/src/trainers/tft/
├── mod.rs (~50 lines) - Module exports and re-exports
├── config.rs (~350 lines) - Configuration structs
├── types.rs (~250 lines) - QAT metrics and helper types
├── model.rs (~100 lines) - TFTModel trait and implementations
├── trainer.rs (~2,000 lines) - Main TFTTrainer implementation
└── tests.rs (~165 lines) - Test module
Current File Structure Analysis
Line Distribution
| Section | Lines | Description |
|---|---|---|
| Header/Imports | 1-33 | Module documentation and use statements |
| QAT Metrics Types | 34-100 | QATMetrics, ScaleStatistics, ZeroPointStatistics, etc. |
| TFTModel Trait | 101-202 | Trait definition and implementations |
| TFTTrainer Struct | 203-277 | Main trainer struct |
| TrainingState | 278-332 | Internal state tracking |
| Progress Types | 333-385 | TrainingProgress, ResourceUsage |
| Config | 386-538 | TFTTrainerConfig struct and impls |
| Trainer Impl | 539-2,498 | Main implementation block with 30+ methods |
| ValidationMetrics | 2,499-2,507 | Helper struct |
| TrainingMetrics | 2,508-2,548 | Public metrics result |
| Tests | 2,550-2,915 | Test module (~365 lines) |
Major Components Identified
-
Configuration Structs (Lines 386-538)
TFTTrainerConfig(main config)impl Default for TFTTrainerConfigimpl TFTTrainerConfig(conversion methods)
-
QAT Metrics Types (Lines 34-100)
QATMetricsScaleStatisticsZeroPointStatisticsObserverRangeStatisticsLayerQuantizationMetrics
-
TFTModel Trait (Lines 101-202)
pub trait TFTModelimpl TFTModel for TemporalFusionTransformer- Commented out QAT implementation
-
Progress/State Types (Lines 278-385)
TrainingState(private)TrainingProgress(public)ResourceUsage(public)
-
Statistics Types (Lines 2,499-2,548)
ValidationMetrics(private)TrainingMetrics(public)
-
Main Trainer (Lines 203-2,498)
TFTTrainerstructimpl Debug for TFTTrainerimpl TFTTrainer(30+ methods)
Proposed Split Plan
File 1: mod.rs (~50 lines)
Purpose: Module organization and re-exports
Contents:
//! TFT (Temporal Fusion Transformer) Trainer Module
//!
//! Production-grade TFT training pipeline with:
//! - GPU acceleration (4GB VRAM optimized)
//! - QAT (Quantization-Aware Training) support
//! - Real-time progress streaming
//! - Checkpoint persistence to MinIO/S3
//!
//! ## Module Structure
//!
//! - `config` - Training configuration and hyperparameters
//! - `types` - QAT metrics and helper types
//! - `model` - TFTModel trait abstraction
//! - `trainer` - Main TFTTrainer implementation
//! - `tests` - Integration tests
mod config;
mod model;
mod trainer;
mod types;
#[cfg(test)]
mod tests;
// Re-export public API
pub use config::TFTTrainerConfig;
pub use model::TFTModel;
pub use trainer::TFTTrainer;
pub use types::{
LayerQuantizationMetrics, ObserverRangeStatistics, QATMetrics, ResourceUsage,
ScaleStatistics, TrainingMetrics, TrainingProgress, ZeroPointStatistics,
};
Lines: ~50
File 2: config.rs (~350 lines)
Purpose: Training configuration and hyperparameters
Contents:
- Lines 386-538:
TFTTrainerConfigstruct - All conversion methods (
to_model_config(),to_training_config()) - Default implementation
- Builder methods (if any)
Key Structs:
pub struct TFTTrainerConfigimpl Default for TFTTrainerConfigimpl TFTTrainerConfig(conversion helpers)
Imports Needed:
use serde::{Deserialize, Serialize};
use crate::checkpoint::CheckpointConfig;
use crate::tft::{TFTConfig, training::TFTTrainingConfig};
Lines: ~350
File 3: types.rs (~250 lines)
Purpose: QAT metrics, progress reporting, and helper types
Contents:
- Lines 34-100: QAT metrics structs
QATMetricsScaleStatisticsZeroPointStatisticsObserverRangeStatisticsLayerQuantizationMetrics
- Lines 333-385: Progress types
TrainingProgressResourceUsage
- Lines 2,508-2,548: Result types
TrainingMetrics
- Internal types:
TrainingState(lines 278-332)ValidationMetrics(lines 2,499-2,507)
Key Structs:
- All
pubmetric/progress types (exported) - Internal state types (not exported, only pub(crate))
Imports Needed:
use std::collections::HashMap;
use serde::{Deserialize, Serialize};
Lines: ~250
File 4: model.rs (~100 lines)
Purpose: TFTModel trait abstraction for polymorphic FP32/QAT support
Contents:
- Lines 101-202: TFTModel trait
pub trait TFTModel: Send + Syncimpl TFTModel for TemporalFusionTransformer- Commented QAT implementation (for future use)
Key Components:
- Trait definition with forward pass, device, config access
- FP32 implementation
- QAT placeholder (currently disabled)
Imports Needed:
use std::sync::Arc;
use candle_core::{Device, Tensor};
use candle_nn::VarMap;
use crate::tft::{TFTConfig, TemporalFusionTransformer};
use crate::MLError;
Lines: ~100
File 5: trainer.rs (~2,000 lines)
Purpose: Main TFTTrainer implementation
Contents:
- Lines 203-277:
TFTTrainerstruct definition - Lines 259-277:
impl Debug for TFTTrainer - Lines 539-2,498:
impl TFTTrainer(all methods)- Constructor:
new() - Public API:
train(),set_progress_callback() - Internal methods: 30+ methods for training, validation, checkpointing, etc.
- Constructor:
Major Methods:
new()- Constructor with auto-batch sizingtrain()- Main training looptrain_epoch()- Single epoch trainingvalidate_epoch()- Validation passbatch_to_tensors()- Data conversioncompute_quantile_loss()- Loss computationcompute_rmse()- RMSE metriccheck_early_stopping()- Early stopping logicsave_checkpoint()- Checkpoint persistencesend_progress_update()- Progress streaming- QAT methods:
run_qat_calibration()qat_to_quantized_checkpoint()export_qat_metrics()apply_qat_lr_schedule()
- Helper methods:
initialize_optimizer()sync_cuda_device()get_resource_usage()- Various getters
Imports Needed:
use std::collections::HashMap;
use std::sync::Arc;
use std::time::{Duration, Instant, SystemTime};
use candle_core::{Device, IndexOp, Tensor};
use candle_nn::VarMap;
use ndarray::Dimension;
use tokio::sync::mpsc;
use tracing::{debug, error, info, instrument, warn};
use crate::checkpoint::{CheckpointConfig, CheckpointManager, CheckpointMetadata, CheckpointStorage};
use crate::memory_optimization::{AutoBatchSizer, BatchSizeConfig, ModelPrecision, OptimizerType};
use crate::tft::training::{TFTBatch, TFTDataLoader};
use crate::{MLError, MLResult};
use super::config::TFTTrainerConfig;
use super::model::TFTModel;
use super::types::{QATMetrics, ResourceUsage, TrainingMetrics, TrainingProgress, TrainingState, ValidationMetrics};
Lines: ~2,000
File 6: tests.rs (~165 lines)
Purpose: Integration tests (separated from main code)
Contents:
- Lines 2,550-2,915: Test module
test_tft_trainer_creation()test_training_config_conversion()test_checkpoint_save_load()- OOM retry tests (multiple tests)
Structure:
use super::*;
use crate::checkpoint::FileSystemStorage;
use std::path::PathBuf;
use tempfile::TempDir;
#[tokio::test]
async fn test_tft_trainer_creation() { ... }
#[tokio::test]
async fn test_training_config_conversion() { ... }
// ... more tests
Lines: ~165
Dependencies Between Modules
Import Graph
mod.rs
├─> config.rs (TFTTrainerConfig)
├─> types.rs (metrics, progress, state)
├─> model.rs (TFTModel trait)
├─> trainer.rs (TFTTrainer)
└─> tests.rs (uses all above)
trainer.rs
├─> config.rs (TFTTrainerConfig)
├─> types.rs (TrainingState, ValidationMetrics, etc.)
└─> model.rs (Box<dyn TFTModel>)
model.rs
├─> crate::tft (TemporalFusionTransformer, TFTConfig)
└─> independent (no internal deps)
config.rs
├─> crate::tft (TFTConfig, TFTTrainingConfig)
└─> independent (no internal deps)
types.rs
└─> independent (no internal deps)
Internal Type Visibility
Private Types (only used in trainer.rs):
TrainingState→ Keep in types.rs aspub(crate)ValidationMetrics→ Keep in types.rs aspub(crate)
Public Types (exported from mod.rs):
- All QAT metrics types
TrainingProgressResourceUsageTrainingMetricsTFTTrainerConfigTFTModeltraitTFTTrainer
Migration Strategy
Step 1: Create Directory Structure
mkdir -p ml/src/trainers/tft
Step 2: Create Files in Order (preserves compilation)
-
Create
types.rs(no dependencies on other split files)- Extract QAT metrics (lines 34-100)
- Extract progress types (lines 333-385)
- Extract TrainingState (lines 278-332)
- Extract ValidationMetrics (lines 2,499-2,507)
- Extract TrainingMetrics (lines 2,508-2,548)
-
Create
model.rs(only depends on external crates)- Extract TFTModel trait (lines 101-202)
-
Create
config.rs(only depends on external crates)- Extract TFTTrainerConfig (lines 386-538)
-
Create
trainer.rs(depends on types, model, config)- Extract TFTTrainer struct (lines 203-277)
- Extract impl blocks (lines 259-277, 539-2,498)
- Update imports to use
super::{types, model, config}
-
Create
tests.rs(depends on all above)- Extract test module (lines 2,550-2,915)
- Remove
#[cfg(test)]attribute (entire file is tests)
-
Create
mod.rs(coordinates all modules)- Module declarations
- Re-exports
-
Update
ml/src/trainers/mod.rs- Change from
pub mod tft;to module with sub-modules - Or keep as
pub mod tft;and let tft/mod.rs handle re-exports
- Change from
Step 3: Update References
Files that import from tft.rs:
- Likely in
ml/src/lib.rsor other trainer files - Should continue to work due to re-exports in
mod.rs
Example Import Changes:
// Before (may still work)
use crate::trainers::tft::TFTTrainer;
// After (explicit, if needed)
use crate::trainers::tft::TFTTrainer; // Same! Re-exported from mod.rs
Step 4: Verification
- Run
cargo checkafter each file creation - Run full test suite:
cargo test --package ml trainers::tft - Run clippy:
cargo clippy --package ml - Format:
cargo fmt --package ml
Risk Assessment
Low Risk
- types.rs: Pure data structures, no complex logic
- config.rs: Configuration only, well-defined boundaries
- model.rs: Trait abstraction, clean separation
Medium Risk
- trainer.rs: Large impl block, many internal methods
- Mitigation: Keep all trainer methods together, don't split impl block
- Mitigation: Use
pub(crate)for internal helpers in types.rs
Potential Issues
-
Circular Dependencies
- Risk: If types.rs needs TFTTrainer, and trainer.rs needs types
- Mitigation: Use
pub(crate)and keep types as pure data structures
-
Import Path Changes
- Risk: External code may need updates
- Mitigation: Re-export everything from mod.rs for backward compatibility
-
Private Method Visibility
- Risk: Internal helpers in TFTTrainer may need visibility adjustments
- Mitigation: Keep all trainer impl in one file, use
pub(crate)sparingly
Comparison with DQN Split
| Aspect | DQN | TFT (Proposed) |
|---|---|---|
| Files | 4 | 6 |
| config.rs | 566 lines | ~350 lines |
| statistics.rs | 134 lines | ~250 lines (types.rs) |
| trainer.rs | 4,007 lines | ~2,000 lines |
| mod.rs | 30 lines | ~50 lines |
| Extra files | - | model.rs (~100), tests.rs (~165) |
Key Differences:
- TFT has TFTModel trait abstraction → separate
model.rs - TFT has extensive QAT metrics → dedicated section in
types.rs - TFT tests are larger → separate
tests.rsfile - DQN has simpler statistics → smaller statistics.rs
Similarities:
- Both split config into dedicated file
- Both extract statistics/types
- Both keep main trainer implementation in trainer.rs
- Both use mod.rs for re-exports
Next Steps
For Implementation Agent:
- Create
ml/src/trainers/tft/directory - Extract files in order (types → model → config → trainer → tests → mod)
- Verify compilation after each file
- Update imports in trainer.rs to use
super::{types, model, config} - Run test suite to ensure no breakage
- Update
ml/src/trainers/mod.rsif needed
Validation Checklist:
- All files compile independently
cargo test --package ml trainers::tftpasses- No new clippy warnings
- Public API unchanged (backward compatible)
- File sizes reasonable (<2,500 lines each)
- No circular dependencies
Conclusion
The TFT trainer can be cleanly split into 6 files following the DQN pattern with minor extensions:
- types.rs replaces statistics.rs (includes QAT metrics)
- model.rs added for trait abstraction
- tests.rs separated due to size
- trainer.rs remains largest but manageable (~2,000 lines vs 2,915)
This split improves: ✅ Modularity: Clear separation of concerns ✅ Maintainability: Smaller, focused files ✅ Testability: Dedicated test file ✅ Compilation Speed: Smaller compilation units ✅ Code Navigation: Easier to find specific components
Estimated total lines after split: 2,915 lines (no change, just reorganized)
Ready for implementation: ✅ Yes, plan is complete and low-risk.