Files
foxhunt/AGENT_136_IMPLEMENTATION_GUIDE.md
jgrusewski 35feadf55e 🚀 Wave 160 Phase 6: CUDA Mandatory + TDD Testing + TFT Complete (21 Agents)
## Major Achievements

### 1. CUDA Made Default & Mandatory (Agent 143)
- CUDA now default feature in ml/Cargo.toml
- All training requires GPU (no silent CPU fallback)
- Added get_training_device() helper with fail-fast errors
- Removed --use-gpu flags (GPU mandatory)
- **Impact**: No more wasting time on accidental CPU training

### 2. TFT Training COMPLETE (Agent 144)
-  Training completed successfully in 7.6 minutes
-  Early stopping at epoch 100/200 (best val loss: 0.097318)
-  11 checkpoints saved to ml/trained_models/production/tft/
-  GPU Performance: 99% utilization, 367MB VRAM, 4.4s/epoch
-  10x speedup vs CPU (4.4s vs 43-55s per epoch)
- **Status**: PRODUCTION READY

### 3. TFT CUDA Tensor Contiguity Fix (Agent 142)
- Fixed "matmul not supported for non-contiguous tensors" error
- Added .contiguous() call after narrow() operation in QuantileLayer
- Enabled CUDA-accelerated TFT training
- **Files**: ml/src/tft/quantile_outputs.rs

### 4. MAMBA-2 CUDA Layer Normalization (Agent 145)
- Created CudaLayerNorm wrapper for missing CUDA kernel
- Implemented manual layer norm: γ * (x - μ) / sqrt(σ² + ε) + β
- MAMBA-2 now runs on CUDA (no more "no cuda implementation" error)
- **Files**: ml/src/mamba/mod.rs

### 5. TDD E2E Test Suite (Agent 146) 
- Created comprehensive MAMBA-2 test suite (297 lines)
- 7 tests: shapes, batches, CUDA, gradients, configs
- **16x faster debugging**: 5s per iteration vs 80s
- Already caught dtype mismatch bug (F32 vs F64)
- **Files**: ml/tests/e2e_mamba2_training.rs

## Agent Summary (Agents 126-146)

### Code Fixes (Parallel - Agents 137-141)
- **Agent 137**: MAMBA-2 batch dimension fix (streaming + batch loaders)
- **Agent 138**: Liquid NN API fix (mutable loader, iterator fix)
- **Agent 139**: PPO CheckpointMetadata fix (signature fields)
- **Agent 140**: Paper trading executor (498 lines, 100ms polling)
- **Agent 141**: Real model loading (RealDQNModel, RealPPOModel)

### Infrastructure (Agents 143-146)
- **Agent 143**: CUDA mandatory (Cargo.toml, device helpers)
- **Agent 144**: TFT verification (completion monitoring)
- **Agent 145**: MAMBA-2 CUDA layer norm wrapper
- **Agent 146**: TDD E2E test suite (16x faster debugging)

## Files Modified

### Core ML Infrastructure
- ml/Cargo.toml: Added default = ["minimal-inference", "cuda"]
- ml/src/lib.rs: Added get_training_device() helper (+109 lines)
- ml/src/tft/quantile_outputs.rs: Fixed tensor contiguity
- ml/src/mamba/mod.rs: Added CudaLayerNorm wrapper (+41 lines)

### Training Scripts
- ml/examples/train_tft_dbn.rs: Removed --use-gpu flag
- ml/examples/train_ppo.rs: Removed --use-gpu flag
- ml/examples/train_mamba2_dbn.rs: Forced CUDA-only mode
- ml/examples/train_liquid_dbn.rs: Fixed API usage

### Data Loaders
- ml/src/data_loaders/dbn_sequence_loader.rs: Fixed batch dimensions
- ml/src/data_loaders/streaming_dbn_loader.rs: Fixed batch dimensions

### Trading Service
- services/trading_service/src/paper_trading_executor.rs: New executor (+498 lines)
- services/trading_service/src/services/enhanced_ml.rs: Real model loading
- services/trading_service/src/ensemble_coordinator.rs: Integration

### Tests
- ml/tests/e2e_mamba2_training.rs: New TDD test suite (+297 lines)

### Trainers
- ml/src/trainers/tft.rs: Fixed CheckpointMetadata signature fields

## Performance Metrics

### TFT Training
- Duration: 7.6 minutes (100 epochs with early stopping)
- GPU Utilization: 99%
- GPU Memory: 367MB / 4GB (9%)
- Epoch Time: 4.4 seconds (vs 43-55s on CPU)
- Speedup: 10x vs CPU
- Status:  PRODUCTION READY

### TDD Testing
- Test Execution: 5-10 seconds per test
- Debugging Iteration: 5 seconds (vs 80 seconds before)
- Speedup: 16x faster debugging
- First Bug Found: <1 minute (dtype mismatch)

## Documentation
- 21 comprehensive agent reports
- TDD quick start guide
- CUDA troubleshooting guide
- Training verification procedures

## Next Steps
1. Fix MAMBA-2 dtype mismatch (F32→F64) - 2 minutes
2. Run MAMBA-2 tests until passing - 5-10 minutes
3. Launch full MAMBA-2 training - 200 epochs
4. Launch Liquid NN training

## System Status
- TFT:  COMPLETE (production ready)
- MAMBA-2: 🧪 IN TESTING (TDD suite ready)
- CUDA:  DEFAULT (mandatory for training)
- Tests:  16x faster debugging

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

Co-Authored-By: Claude <noreply@anthropic.com>
2025-10-14 23:13:34 +02:00

632 lines
18 KiB
Markdown

# Agent 136: Model Loading Implementation Guide
**For**: Next developer implementing real model loading
**Priority**: CRITICAL (blocks paper trading)
**Estimated Effort**: 7-11 hours
---
## QUICK START
### Problem
Trading service uses `MockMLModelWrapper` instead of loading real trained models from safetensors checkpoints.
### Solution
Replace mock implementations with actual model loading using the `ml` crate's checkpoint infrastructure.
---
## IMPLEMENTATION STEPS
### Step 1: Add Model Loading Function (2-3 hours)
**File**: `services/trading_service/src/services/enhanced_ml.rs`
**Replace lines 210-244** with:
```rust
use candle_core::{Device, DType};
use candle_nn::VarBuilder;
use ml::dqn::DQNAgent;
use ml::ppo::PPOAgent;
use std::path::Path;
async fn load_model_from_file(
&self,
model_id: &str,
checkpoint_path: &Path,
) -> Result<Arc<dyn MLModel>, String> {
info!("Loading model {} from {}", model_id, checkpoint_path.display());
// Check file exists
if !checkpoint_path.exists() {
return Err(format!("Checkpoint not found: {}", checkpoint_path.display()));
}
// Select device (GPU if available, CPU fallback)
let device = Device::cuda_if_available(0)
.map_err(|e| format!("Failed to initialize device: {}", e))?;
info!("Using device: {:?}", device);
// Detect model type from model_id
let model_type = if model_id.contains("DQN") {
ModelType::DQN
} else if model_id.contains("PPO") {
ModelType::PPO
} else if model_id.contains("TFT") {
ModelType::TFT
} else if model_id.contains("MAMBA") {
ModelType::MAMBA2
} else {
return Err(format!("Unknown model type in model_id: {}", model_id));
};
// Load model based on type
let model: Arc<dyn MLModel> = match model_type {
ModelType::DQN => {
// Load DQN from safetensors
let config = ml::dqn::DQNConfig {
state_dim: 16, // From paper_trading_config.yaml
action_dim: 3, // Buy/Sell/Hold
hidden_dim: 256,
learning_rate: 0.0001,
gamma: 0.99,
epsilon_start: 0.1, // Low epsilon for production
epsilon_end: 0.01,
epsilon_decay: 0.995,
replay_buffer_size: 100000,
batch_size: 128,
};
let mut agent = DQNAgent::new(config, device.clone())
.map_err(|e| format!("Failed to create DQN agent: {}", e))?;
// Load checkpoint weights
agent.load_checkpoint(checkpoint_path)
.map_err(|e| format!("Failed to load DQN checkpoint: {}", e))?;
Arc::new(agent) as Arc<dyn MLModel>
}
ModelType::PPO => {
// Load PPO from safetensors (actor + critic)
let config = ml::ppo::PPOConfig {
state_dim: 16,
action_dim: 3,
hidden_dim: 256,
learning_rate: 0.0003,
gamma: 0.99,
gae_lambda: 0.95,
clip_epsilon: 0.2,
value_clip: 0.2,
entropy_coeff: 0.01,
max_grad_norm: 0.5,
batch_size: 64,
epochs_per_update: 10,
};
let mut agent = PPOAgent::new(config, device.clone())
.map_err(|e| format!("Failed to create PPO agent: {}", e))?;
// PPO has separate actor/critic checkpoints
// Parse checkpoint paths from model_id or use convention
let actor_path = checkpoint_path.parent()
.ok_or("Invalid checkpoint path")?
.join(format!("{}_actor.safetensors", model_id));
let critic_path = checkpoint_path.parent()
.ok_or("Invalid checkpoint path")?
.join(format!("{}_critic.safetensors", model_id));
agent.load_checkpoint(&actor_path, &critic_path)
.map_err(|e| format!("Failed to load PPO checkpoint: {}", e))?;
Arc::new(agent) as Arc<dyn MLModel>
}
_ => {
return Err(format!("Model type {:?} not yet implemented for loading", model_type));
}
};
info!("Successfully loaded model {} ({})",
model_id,
format_size(checkpoint_path.metadata()
.map(|m| m.len())
.unwrap_or(0)));
Ok(model)
}
// Helper to format file size
fn format_size(bytes: u64) -> String {
if bytes < 1024 {
format!("{}B", bytes)
} else if bytes < 1024 * 1024 {
format!("{:.1}KB", bytes as f64 / 1024.0)
} else {
format!("{:.1}MB", bytes as f64 / 1024.0 / 1024.0)
}
}
```
---
### Step 2: Update Ensemble Coordinator (1-2 hours)
**File**: `services/trading_service/src/ensemble_coordinator.rs`
**A. Store Loaded Models in Registry**
Update `ModelRegistry` struct (line 214):
```rust
pub struct ModelRegistry {
/// Active models (currently serving predictions)
active: HashMap<String, Arc<dyn MLModel>>, // Changed from String to Arc<dyn MLModel>
/// Shadow models (staged for hot-swap)
shadow: HashMap<String, Arc<dyn MLModel>>,
}
```
**B. Add Model Registration Method**
Add to `EnsembleCoordinator` (after line 87):
```rust
/// Register a loaded model in the ensemble
pub async fn register_loaded_model(
&self,
model_id: String,
model: Arc<dyn MLModel>,
weight: f64,
) -> MLResult<()> {
// Register weight
let model_weight = ModelWeight::new(model_id.clone(), weight);
let mut weights = self.model_weights.write().await;
weights.insert(model_id.clone(), model_weight);
// Store model in registry
let mut registry = self.active_models.write().await;
registry.active.insert(model_id.clone(), model);
info!("Registered model {} with weight {} (model loaded)", model_id, weight);
Ok(())
}
```
**C. Replace Mock Predictions**
Replace `generate_mock_predictions` (lines 130-169) with:
```rust
async fn generate_real_predictions(
&self,
features: &Features,
) -> MLResult<Vec<ModelPrediction>> {
let registry = self.active_models.read().await;
let weights = self.model_weights.read().await;
let mut predictions = Vec::new();
for (model_id, model) in registry.active.iter() {
if !weights.contains_key(model_id) {
warn!("Model {} in registry but not in weights, skipping", model_id);
continue;
}
// Real model inference
match model.predict(features).await {
Ok(prediction) => {
debug!("Model {} predicted: value={:.3}, confidence={:.3}",
model_id, prediction.value, prediction.confidence);
predictions.push(prediction);
}
Err(e) => {
warn!("Model {} prediction failed: {}", model_id, e);
// Continue with other models (ensemble degradation handling)
}
}
}
if predictions.is_empty() {
return Err(MLError::InferenceError(
"No successful predictions from any model".to_string()
));
}
Ok(predictions)
}
```
**D. Update predict() Method**
Update line 100 to call real predictions:
```rust
pub async fn predict(&self, features: &Features) -> MLResult<EnsembleDecision> {
debug!("Making ensemble prediction with {} features", features.values.len());
let start_time = Instant::now();
// Real model predictions
let predictions = self.generate_real_predictions(features).await?;
// Rest of method unchanged...
let decision = self.aggregator.aggregate(
predictions,
&*self.model_weights.read().await,
).await?;
let aggregation_latency_us = start_time.elapsed().as_micros() as f64;
info!(
"Ensemble decision: {:?}, confidence: {:.3}, disagreement: {:.3}, latency: {:.1}μs",
decision.action, decision.confidence, decision.disagreement_rate, aggregation_latency_us
);
Ok(decision)
}
```
---
### Step 3: Initialize Models on Startup (2-3 hours)
**File**: `services/trading_service/src/main.rs`
**A. Add Config Loading**
Add to imports:
```rust
use serde::{Deserialize, Serialize};
use std::fs;
```
Add config structs:
```rust
#[derive(Debug, Deserialize)]
struct PaperTradingConfig {
ensemble: EnsembleConfig,
}
#[derive(Debug, Deserialize)]
struct EnsembleConfig {
models: Vec<ModelConfig>,
}
#[derive(Debug, Deserialize)]
struct ModelConfig {
name: String,
#[serde(rename = "type")]
model_type: String,
checkpoint: Option<String>,
checkpoint_actor: Option<String>,
checkpoint_critic: Option<String>,
weight: f64,
enabled: bool,
}
```
**B. Add Model Initialization Function**
```rust
async fn initialize_ensemble_models(
state: &TradingServiceState,
) -> Result<(), Box<dyn std::error::Error>> {
info!("Initializing ensemble models from config...");
// Load paper trading config
let config_path = "config/paper_trading_config.yaml";
let config_str = fs::read_to_string(config_path)
.context(format!("Failed to read config: {}", config_path))?;
let config: PaperTradingConfig = serde_yaml::from_str(&config_str)
.context("Failed to parse paper trading config")?;
info!("Found {} models in config", config.ensemble.models.len());
// Load each model
let mut loaded_count = 0;
for model_config in &config.ensemble.models {
if !model_config.enabled {
info!("Skipping disabled model: {}", model_config.name);
continue;
}
info!("Loading model: {} (type: {}, weight: {})",
model_config.name, model_config.model_type, model_config.weight);
let checkpoint_path = match model_config.checkpoint {
Some(ref path) => PathBuf::from(path),
None => {
warn!("Model {} has no checkpoint path, skipping", model_config.name);
continue;
}
};
// Load model using EnhancedMLServiceImpl
match state.ml_service.load_model_from_file(&model_config.name, &checkpoint_path).await {
Ok(model) => {
// Register model in ensemble
if let Some(ref coordinator) = state.ensemble_coordinator {
coordinator.register_loaded_model(
model_config.name.clone(),
model,
model_config.weight,
).await?;
loaded_count += 1;
info!("✓ Model {} loaded and registered", model_config.name);
} else {
warn!("No ensemble coordinator available");
}
}
Err(e) => {
warn!("Failed to load model {}: {}", model_config.name, e);
// Continue with other models (allow partial ensemble)
}
}
}
info!("Ensemble initialization complete: {}/{} models loaded",
loaded_count, config.ensemble.models.len());
if loaded_count == 0 {
return Err("No models loaded successfully".into());
}
Ok(())
}
```
**C. Call Initialization in main()**
Add after service state creation (around line where `TradingServiceState` is created):
```rust
// Initialize service state
let state = TradingServiceState::new(...).await?;
// Load ensemble models from paper trading config
initialize_ensemble_models(&state).await?;
info!("Trading service ready with {} models",
state.ensemble_coordinator
.as_ref()
.map(|c| c.model_count())
.unwrap_or(0));
// Start gRPC server
// ...
```
---
### Step 4: Add Unit Tests (2-3 hours)
**File**: `services/trading_service/tests/model_loading_test.rs` (NEW)
```rust
use std::path::PathBuf;
use trading_service::services::enhanced_ml::EnhancedMLServiceImpl;
#[tokio::test]
async fn test_dqn_checkpoint_loading() {
let checkpoint_path = PathBuf::from("ml/trained_models/production/dqn/dqn_epoch_30.safetensors");
assert!(
checkpoint_path.exists(),
"DQN checkpoint not found at {}. Run training first: cargo run -p ml --example train_dqn",
checkpoint_path.display()
);
let ml_service = EnhancedMLServiceImpl::new();
let model = ml_service
.load_model_from_file("DQN_epoch30", &checkpoint_path)
.await
.expect("Failed to load DQN checkpoint");
// Test inference
let features = ml::Features::new(vec![0.5; 16]);
let prediction = model.predict(&features).await.expect("Prediction failed");
assert!(prediction.value >= 0.0 && prediction.value <= 1.0,
"Prediction value out of range: {}", prediction.value);
assert!(prediction.confidence >= 0.0 && prediction.confidence <= 1.0,
"Confidence out of range: {}", prediction.confidence);
println!("✓ DQN checkpoint loaded successfully");
println!(" Prediction: {:.3} (confidence: {:.3})", prediction.value, prediction.confidence);
}
#[tokio::test]
async fn test_ppo_checkpoint_loading() {
let actor_path = PathBuf::from("ml/trained_models/production/ppo/ppo_actor_epoch_130.safetensors");
let critic_path = PathBuf::from("ml/trained_models/production/ppo/ppo_critic_epoch_130.safetensors");
assert!(actor_path.exists(), "PPO actor checkpoint not found");
assert!(critic_path.exists(), "PPO critic checkpoint not found");
let ml_service = EnhancedMLServiceImpl::new();
let model = ml_service
.load_model_from_file("PPO_epoch130", &actor_path)
.await
.expect("Failed to load PPO checkpoint");
let features = ml::Features::new(vec![0.5; 16]);
let prediction = model.predict(&features).await.expect("Prediction failed");
assert!(prediction.value >= 0.0 && prediction.value <= 1.0);
assert!(prediction.confidence >= 0.0 && prediction.confidence <= 1.0);
println!("✓ PPO checkpoint loaded successfully");
}
#[tokio::test]
async fn test_ensemble_with_real_models() {
use trading_service::ensemble_coordinator::EnsembleCoordinator;
let coordinator = EnsembleCoordinator::new();
let ml_service = EnhancedMLServiceImpl::new();
// Load DQN
let dqn_path = PathBuf::from("ml/trained_models/production/dqn/dqn_epoch_30.safetensors");
let dqn_model = ml_service.load_model_from_file("DQN", &dqn_path).await.unwrap();
coordinator.register_loaded_model("DQN".to_string(), dqn_model, 0.4).await.unwrap();
// Load PPO
let ppo_path = PathBuf::from("ml/trained_models/production/ppo/ppo_actor_epoch_130.safetensors");
let ppo_model = ml_service.load_model_from_file("PPO", &ppo_path).await.unwrap();
coordinator.register_loaded_model("PPO".to_string(), ppo_model, 0.6).await.unwrap();
// Verify model count
assert_eq!(coordinator.model_count().await, 2, "Should have 2 models registered");
// Test ensemble prediction
let features = ml::Features::new(vec![0.5; 16]);
let decision = coordinator.predict(&features).await.expect("Ensemble prediction failed");
println!("✓ Ensemble prediction successful");
println!(" Action: {:?}", decision.action);
println!(" Confidence: {:.3}", decision.confidence);
println!(" Disagreement: {:.3}", decision.disagreement_rate);
// Verify prediction is not default/mock
assert!(decision.confidence > 0.0, "Confidence should be > 0");
}
```
---
## VERIFICATION CHECKLIST
After implementation, verify:
### 1. Build Success
```bash
cd services/trading_service
cargo build --release
```
### 2. Unit Tests Pass
```bash
cargo test --package trading_service model_loading
```
### 3. Service Starts Successfully
```bash
cargo run --release
```
**Expected Logs**:
```
[INFO] Initializing ensemble models from config...
[INFO] Found 3 models in config
[INFO] Loading model: DQN_epoch30 (type: DQN, weight: 0.4)
[INFO] Using device: Cuda(0)
[INFO] Successfully loaded model DQN_epoch30 (74.0KB)
[INFO] ✓ Model DQN_epoch30 loaded and registered
[INFO] Loading model: PPO_epoch130 (type: PPO, weight: 0.4)
[INFO] Successfully loaded model PPO_epoch130 (84.0KB)
[INFO] ✓ Model PPO_epoch130 loaded and registered
[INFO] Ensemble initialization complete: 3/3 models loaded
[INFO] Trading service ready with 3 models
```
### 4. Ensemble Predictions Work
```bash
# Use tli to test prediction
tli predict --symbol ES.FUT --features 0.5,0.5,0.5,0.5,0.5,0.5,0.5,0.5,0.5,0.5,0.5,0.5,0.5,0.5,0.5,0.5
```
**Expected Output**:
```
Ensemble Decision:
Action: Buy
Confidence: 0.782
Disagreement: 0.145
Latency: 23.4μs
```
### 5. Check Metrics
```bash
curl http://localhost:9092/metrics | grep ensemble
```
**Expected Metrics**:
```
ensemble_prediction_confidence{symbol="ES.FUT"} 0.782
ensemble_prediction_disagreement{symbol="ES.FUT"} 0.145
ensemble_aggregation_latency_us{symbol="ES.FUT"} 23.4
```
---
## TROUBLESHOOTING
### Error: "Checkpoint not found"
**Fix**: Verify checkpoint paths are relative to project root:
```bash
ls -la ml/trained_models/production/dqn/
ls -la ml/trained_models/production/ppo/
```
### Error: "Failed to initialize device"
**Fix**: Check CUDA availability:
```bash
nvidia-smi
export CUDA_VISIBLE_DEVICES=0
```
### Error: "Failed to load DQN checkpoint: dimension mismatch"
**Fix**: Check model config dimensions match checkpoint:
```rust
// Checkpoint was trained with 16 features, 3 actions
state_dim: 16, // Must match training config
action_dim: 3, // Buy/Sell/Hold
```
### Error: "No successful predictions from any model"
**Fix**: Check model logs for individual failures:
```bash
docker-compose logs trading_service | grep -i "prediction failed"
```
---
## DEPENDENCIES
Add to `services/trading_service/Cargo.toml`:
```toml
[dependencies]
ml = { path = "../../ml" }
candle-core = "0.7"
candle-nn = "0.7"
serde_yaml = "0.9"
anyhow = "1.0"
```
---
## ESTIMATED TIMELINE
- **Step 1** (Model Loading): 2-3 hours
- **Step 2** (Ensemble Update): 1-2 hours
- **Step 3** (Startup Init): 2-3 hours
- **Step 4** (Unit Tests): 2-3 hours
- **Testing & Debug**: 1-2 hours
**Total**: 7-11 hours (1-2 business days)
---
## SUCCESS CRITERIA
✅ All unit tests pass
✅ Service starts without errors
✅ Logs show "Successfully loaded model" for all 3 models
✅ Ensemble produces non-zero predictions
✅ Prediction confidence > 0.55 (trading threshold)
✅ Latency < 50μs P99
✅ Paper trading generates orders (not 0)
---
## NEXT STEPS AFTER COMPLETION
1. Monitor paper trading for 24 hours
2. Verify orders are being generated
3. Check Sharpe ratio matches expected (~1.5-1.6)
4. Validate disagreement rates (~10-30%)
5. Proceed to Phase 2: 1% capital deployment