# TFT Trainer Extraction Guide **For Implementation Agent** This document provides exact line ranges and extraction instructions for splitting `ml/src/trainers/tft.rs`. --- ## Extraction Order & Line Ranges ### 1. Create `types.rs` (~250 lines) Extract the following sections in order: #### Section 1: QAT Metrics (Lines 34-100) ```rust // ============================================================================ // QAT Metrics Export Types (Prometheus-compatible) // ============================================================================ /// Comprehensive QAT metrics for Prometheus/Grafana export #[derive(Debug, Clone, Serialize, Deserialize)] pub struct QATMetrics { ... } #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ScaleStatistics { ... } #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ZeroPointStatistics { ... } #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ObserverRangeStatistics { ... } #[derive(Debug, Clone, Serialize, Deserialize)] pub struct LayerQuantizationMetrics { ... } ``` **Lines to extract**: 34-100 (67 lines) #### Section 2: Internal State (Lines 278-332) ```rust /// Training state tracking #[derive(Debug, Clone)] pub(crate) struct TrainingState { ... } // Change to pub(crate) impl Default for TrainingState { ... } ``` **Lines to extract**: 278-332 (55 lines) **Important**: Change `struct TrainingState` from no visibility to `pub(crate)` #### Section 3: Progress Types (Lines 333-385) ```rust /// Training progress update for gRPC streaming #[derive(Debug, Clone, Serialize, Deserialize)] pub struct TrainingProgress { ... } /// Resource usage statistics #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ResourceUsage { ... } impl Default for ResourceUsage { ... } ``` **Lines to extract**: 333-385 (53 lines) #### Section 4: Internal Metrics (Lines 2,499-2,507) ```rust /// Validation metrics #[derive(Debug, Clone, Default)] pub(crate) struct ValidationMetrics { ... } // Change to pub(crate) ``` **Lines to extract**: 2,499-2,507 (9 lines) **Important**: Change `struct ValidationMetrics` visibility to `pub(crate)` #### Section 5: Result Metrics (Lines 2,508-2,548) ```rust /// Training metrics result #[derive(Debug, Clone, Default, Serialize, Deserialize)] pub struct TrainingMetrics { ... } ``` **Lines to extract**: 2,508-2,548 (41 lines) #### types.rs Header ```rust //! TFT Trainer Types //! //! Data structures for metrics, progress reporting, and internal state. use std::collections::HashMap; use serde::{Deserialize, Serialize}; // Re-import internal types if needed use std::time::Instant; ``` **Total lines**: ~250 --- ### 2. Create `model.rs` (~100 lines) Extract lines 101-202: ```rust //! TFT Model Trait Abstraction //! //! Polymorphic interface for FP32 and QAT TFT models. use std::sync::Arc; use candle_core::{Device, Tensor}; use candle_nn::VarMap; use crate::tft::{TFTConfig, TemporalFusionTransformer}; use crate::MLError; /// Trait for polymorphic TFT model (FP32 or QAT) pub trait TFTModel: Send + Sync { ... } /// Implement TFTModel for standard FP32 TemporalFusionTransformer impl TFTModel for TemporalFusionTransformer { ... } // Implement TFTModel for QAT TemporalFusionTransformer - DISABLED /* QAT IMPLEMENTATION DISABLED DUE TO P0 COMPILATION ERRORS impl TFTModel for QATTemporalFusionTransformer { ... } */ ``` **Lines to extract**: 101-202 (102 lines) --- ### 3. Create `config.rs` (~350 lines) Extract lines 386-538: ```rust //! TFT Trainer Configuration //! //! Training hyperparameters and configuration management. use serde::{Deserialize, Serialize}; use crate::checkpoint::CheckpointConfig; use crate::tft::{TFTConfig, training::TFTTrainingConfig}; /// TFT trainer configuration from gRPC proto #[derive(Debug, Clone, Serialize, Deserialize)] pub struct TFTTrainerConfig { ... } impl Default for TFTTrainerConfig { ... } impl TFTTrainerConfig { pub fn to_model_config(&self) -> TFTConfig { ... } pub fn to_training_config(&self) -> TFTTrainingConfig { ... } } ``` **Lines to extract**: 386-538 (153 lines) --- ### 4. Create `trainer.rs` (~2,000 lines) This is the largest file. Extract in sections: #### Section 1: Header (Lines 1-33) Use the module documentation but update: ```rust //! Temporal Fusion Transformer Trainer Implementation //! //! Production-grade TFT trainer optimized for GPU (4GB VRAM) with checkpoint //! management, real-time metrics reporting, and MinIO/S3 storage integration. //! //! ## Features //! - GPU acceleration with memory-efficient attention //! - Quantile loss for probabilistic forecasting //! - Attention weights analysis //! - Real-time training progress streaming //! - Checkpoint persistence to MinIO/S3 //! - RMSE and quantile loss metrics 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, TFTTrainingConfig}; use crate::tft::{TFTConfig, TemporalFusionTransformer}; use crate::{MLError, MLResult}; // Import from sibling modules use super::config::TFTTrainerConfig; use super::model::TFTModel; use super::types::{ QATMetrics, ResourceUsage, TrainingMetrics, TrainingProgress, TrainingState, ValidationMetrics, }; ``` #### Section 2: TFTTrainer Struct (Lines 203-277) ```rust /// TFT trainer with gRPC interface integration pub struct TFTTrainer { ... } impl std::fmt::Debug for TFTTrainer { ... } ``` **Lines to extract**: 203-277 (75 lines) #### Section 3: Full Implementation (Lines 539-2,498) ```rust impl TFTTrainer { /// Create new TFT trainer instance pub fn new(...) -> MLResult { ... } pub fn set_progress_callback(...) { ... } fn initialize_optimizer(&mut self) -> MLResult<()> { ... } fn sync_cuda_device(device: &Device) -> MLResult<()> { ... } fn recreate_data_loader_with_batch_size(...) { ... } pub async fn train(...) -> MLResult { ... } async fn train_epoch(...) -> MLResult { ... } async fn validate_epoch(...) -> MLResult { ... } fn batch_to_tensors(...) -> MLResult<(Tensor, Tensor, Tensor, Tensor)> { ... } fn compute_quantile_loss(...) -> MLResult { ... } fn compute_rmse(...) -> MLResult { ... } fn extract_attention_entropy(&self) -> MLResult> { ... } fn check_early_stopping(&mut self, val_loss: f64) -> bool { ... } async fn save_checkpoint(...) -> MLResult<()> { ... } async fn send_progress_update(...) { ... } async fn send_qat_calibration_progress(&self) { ... } fn get_resource_usage(&self) -> ResourceUsage { ... } pub fn get_model(&self) -> &dyn TFTModel { ... } pub fn get_varmap(&self) -> Arc { ... } pub fn get_training_config(&self) -> &TFTTrainingConfig { ... } pub fn get_qat_min_batch_size(&self) -> usize { ... } pub fn update_batch_size(&mut self, new_batch_size: usize) { ... } async fn quantize_and_save_int8_checkpoint(...) -> MLResult<()> { ... } async fn run_qat_calibration(...) -> MLResult<()> { ... } async fn qat_to_quantized_checkpoint(...) -> MLResult<()> { ... } fn export_qat_metrics(&self) -> Option { ... } fn apply_qat_lr_schedule(&mut self, epoch: usize) -> MLResult<()> { ... } } ``` **Lines to extract**: 539-2,498 (1,960 lines) **Total trainer.rs**: ~2,000 lines --- ### 5. Create `tests.rs` (~165 lines) Extract lines 2,550-2,915: ```rust //! TFT Trainer Integration Tests 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() { ... } #[tokio::test] async fn test_checkpoint_save_load() { ... } #[tokio::test] async fn test_oom_retry_exponential_backoff() { ... } #[tokio::test] async fn test_oom_retry_minimum_batch_size() { ... } ``` **Lines to extract**: 2,550-2,915 (365 lines) **Important**: Remove the `#[cfg(test)]` attribute (entire file is tests) --- ### 6. Create `mod.rs` (~50 lines) New file (not extracted): ```rust //! 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 via gRPC //! - Checkpoint persistence to MinIO/S3 //! - Comprehensive metrics and monitoring //! //! ## Module Structure //! //! - `config` - Training configuration and hyperparameters //! - `types` - QAT metrics, progress reporting, and internal state //! - `model` - TFTModel trait for polymorphic FP32/QAT support //! - `trainer` - Main TFTTrainer implementation //! - `tests` - Integration tests //! //! ## Example Usage //! //! ```rust,no_run //! use foxhunt_ml::trainers::tft::{TFTTrainer, TFTTrainerConfig}; //! use std::sync::Arc; //! //! # async fn example() -> Result<(), Box> { //! let config = TFTTrainerConfig::default(); //! let storage = Arc::new(/* checkpoint storage */); //! let trainer = TFTTrainer::new(config, storage)?; //! let metrics = trainer.train(train_loader, val_loader).await?; //! # Ok(()) //! # } //! ``` 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, }; ``` --- ## Visibility Changes Summary | Type | Original | New | Reason | |------|----------|-----|--------| | `TrainingState` | (no modifier) | `pub(crate)` | Used by trainer.rs | | `ValidationMetrics` | (no modifier) | `pub(crate)` | Used by trainer.rs | | All QAT metrics | `pub` | `pub` | Public API | | `TrainingProgress` | `pub` | `pub` | Public API | | `ResourceUsage` | `pub` | `pub` | Public API | | `TrainingMetrics` | `pub` | `pub` | Public API | | `TFTModel` | `pub` | `pub` | Public API | | `TFTTrainerConfig` | `pub` | `pub` | Public API | | `TFTTrainer` | `pub` | `pub` | Public API | --- ## Import Updates ### In `trainer.rs` **Add these imports at the top:** ```rust // Import from sibling modules use super::config::TFTTrainerConfig; use super::model::TFTModel; use super::types::{ QATMetrics, ResourceUsage, TrainingMetrics, TrainingProgress, TrainingState, ValidationMetrics, }; ``` **Remove these from external imports** (now local): - `TFTTrainerConfig` (now from `super::config`) - Any references to types now in `types.rs` ### In `tests.rs` **Add at the top:** ```rust use super::*; // Imports everything from parent module use crate::checkpoint::FileSystemStorage; use std::path::PathBuf; use tempfile::TempDir; ``` --- ## Compilation Verification Steps After each file creation: ```bash # 1. After creating types.rs cargo check --package ml # 2. After creating model.rs cargo check --package ml # 3. After creating config.rs cargo check --package ml # 4. After creating trainer.rs cargo check --package ml # 5. After creating tests.rs cargo check --package ml # 6. After creating mod.rs and removing old tft.rs cargo check --package ml cargo test --package ml trainers::tft cargo clippy --package ml ``` --- ## Post-Migration Checklist - [ ] All 6 files created in `ml/src/trainers/tft/` - [ ] Original `tft.rs` removed or renamed - [ ] `cargo check` passes - [ ] `cargo test` passes - [ ] `cargo clippy` shows no new warnings - [ ] Public API unchanged (verify with external imports) - [ ] Re-exports in `mod.rs` match public API - [ ] All visibility modifiers correct (`pub`, `pub(crate)`) - [ ] No circular dependencies - [ ] File sizes reasonable: - [ ] mod.rs < 100 lines - [ ] config.rs < 500 lines - [ ] types.rs < 300 lines - [ ] model.rs < 150 lines - [ ] trainer.rs < 2,500 lines - [ ] tests.rs < 400 lines --- ## Troubleshooting ### Issue: "cannot find type `TrainingState` in this scope" **Solution**: Import with `use super::types::TrainingState;` in trainer.rs ### Issue: "type `TrainingState` is private" **Solution**: Change to `pub(crate) struct TrainingState` in types.rs ### Issue: "circular dependency detected" **Solution**: Ensure types.rs doesn't import from trainer.rs ### Issue: "trait `TFTModel` not in scope" **Solution**: Import with `use super::model::TFTModel;` in trainer.rs ### Issue: Tests fail to compile **Solution**: Ensure `use super::*;` is at top of tests.rs --- ## Final File Structure Verification ``` ml/src/trainers/tft/ ├── mod.rs (✓ exports match public API) ├── config.rs (✓ TFTTrainerConfig) ├── types.rs (✓ all metrics/progress types) ├── model.rs (✓ TFTModel trait) ├── trainer.rs (✓ TFTTrainer impl) └── tests.rs (✓ all tests) ``` **Ready to implement!** 🚀