Files
foxhunt/docs/tft-split-extraction-guide.md
jgrusewski 2df1ea92e1 feat(ml): WAVE 29 DQN Codebase Cleanup & Refactoring Campaign
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>
2025-11-27 23:46:13 +01:00

511 lines
13 KiB
Markdown

# 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<Self> { ... }
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<TrainingMetrics> { ... }
async fn train_epoch(...) -> MLResult<f64> { ... }
async fn validate_epoch(...) -> MLResult<ValidationMetrics> { ... }
fn batch_to_tensors(...) -> MLResult<(Tensor, Tensor, Tensor, Tensor)> { ... }
fn compute_quantile_loss(...) -> MLResult<Tensor> { ... }
fn compute_rmse(...) -> MLResult<f64> { ... }
fn extract_attention_entropy(&self) -> MLResult<Option<f64>> { ... }
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<VarMap> { ... }
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<QATMetrics> { ... }
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<dyn std::error::Error>> {
//! 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!** 🚀