Files
foxhunt/ml/tests/real_data_helpers.rs
jgrusewski aac0597cd2 feat(ml): DQN Option B checkpoint fix + TFT OOM investigation
- Fixed DQN early stopping checkpoint naming bug (Option B)
  - Added is_final: bool parameter to checkpoint callback signature
  - Trainer now distinguishes final checkpoints from regular epoch checkpoints
  - Final checkpoints use 'dqn_final_epoch{N}' naming convention
  - Regular checkpoints use 'dqn_epoch_{N}' naming convention

- Completed comprehensive TFT OOM investigation
  - Spawned 3 parallel agents for memory analysis
  - Identified 16.4GB memory leak (29.7x over expected 525-550MB)
  - Root causes: Attention cache bloat (960MB), gradient accumulation bug, detached tensors
  - Recommended fixes: Disable cache during training, explicit tensor drops
  - Created TFT_MEMORY_ANALYSIS.md, TFT_MEMORY_LEAK_ANALYSIS.md

- DQN 100-epoch training VERIFIED on Runpod RTX A4000
  - Training completed successfully: 100/100 epochs
  - Final checkpoint created: dqn_final_epoch100.safetensors
  - Training speed: 4.8 sec/epoch (3.5x faster than baseline)
  - Option B fix working perfectly

- Deployed RTX 4090 pod for TFT testing
  - Pod ID: 6244yzm9hadnog
  - 24GB VRAM to bypass OOM issue
  - EUR-IS-1 datacenter, $0.59/hr

Files modified:
- ml/examples/train_dqn.rs (checkpoint callback signature)
- ml/src/trainers/dqn.rs (callback signature + is_final parameter)
- CLAUDE.md (compacted to ~11k chars)

Generated reports:
- TFT_MEMORY_ANALYSIS.md (15-section memory breakdown)
- TFT_MEMORY_QUICK_SUMMARY.md (executive summary)
- TFT_MEMORY_LEAK_ANALYSIS.md (5 critical leaks identified)

Co-Authored-By: Claude <noreply@anthropic.com>
2025-10-25 23:49:24 +02:00

300 lines
9.6 KiB
Rust

//! Real Market Data Helpers for ML Model Unit Tests
//!
//! Provides utilities to load real BTC/ETH Parquet data and convert it to formats
//! used by ML model tests (MAMBA-2, DQN, PPO, TFT).
//!
//! # Architecture
//!
//! - Loads real market data from test_data/real/parquet/*.parquet files
//! - Converts to model-specific formats (states, features, time series)
//! - Provides realistic data distributions for proper model testing
//!
//! # Usage
//!
//! ```ignore
//! use real_data_helpers::{RealDataLoader, load_dqn_states, load_tft_sequences};
//!
//! // DQN/PPO tests
//! let states = load_dqn_states(100).await?;
//!
//! // TFT/MAMBA-2 tests
//! let sequences = load_tft_sequences(50, 20).await?;
//! ```
use anyhow::{Context, Result};
use data::parquet_persistence::{MarketDataEvent, ParquetMarketDataReader};
use std::path::PathBuf;
/// Path to real test data directory (relative to workspace root)
const REAL_DATA_PATH: &str = "test_data/real/parquet";
const BTC_FILE: &str = "BTC-USD_30day_2024-09.parquet";
const ETH_FILE: &str = "ETH-USD_30day_2024-09.parquet";
/// Real market data loader for ML tests
pub(crate) struct RealDataLoader {
base_path: String,
}
impl RealDataLoader {
/// Create new loader with workspace-relative path
pub(crate) fn new() -> Self {
let manifest_dir = env!("CARGO_MANIFEST_DIR");
let base_path = PathBuf::from(manifest_dir)
.parent()
.expect("Failed to get workspace root")
.join(REAL_DATA_PATH);
Self {
base_path: base_path.to_string_lossy().to_string(),
}
}
/// Check if real data files exist
pub(crate) fn files_exist(&self) -> bool {
let btc_path = PathBuf::from(&self.base_path).join(BTC_FILE);
let eth_path = PathBuf::from(&self.base_path).join(ETH_FILE);
btc_path.exists() && eth_path.exists()
}
/// Load raw market events from BTC data
pub(crate) async fn load_btc_events(&self, count: usize) -> Result<Vec<MarketDataEvent>> {
self.load_events(BTC_FILE, count).await
}
/// Load raw market events from ETH data
#[allow(dead_code)]
pub(crate) async fn load_eth_events(&self, count: usize) -> Result<Vec<MarketDataEvent>> {
self.load_events(ETH_FILE, count).await
}
/// Load events from a specific file
async fn load_events(&self, filename: &str, count: usize) -> Result<Vec<MarketDataEvent>> {
let reader = ParquetMarketDataReader::new(self.base_path.clone());
let mut events = reader
.read_file(filename)
.await
.with_context(|| format!("Failed to load {}", filename))?;
// Take only the requested number of events
events.truncate(count);
Ok(events)
}
}
impl Default for RealDataLoader {
fn default() -> Self {
Self::new()
}
}
/// Convert market events to DQN/PPO state vectors
///
/// Each state vector contains [price, volume, bid_ask_spread, momentum, volatility]
pub(crate) fn events_to_dqn_states(events: &[MarketDataEvent], state_dim: usize) -> Vec<Vec<f32>> {
if events.is_empty() {
return vec![];
}
events
.windows(5) // Need 5 events to calculate momentum/volatility
.map(|window| {
let current = &window[4];
let prev = &window[0];
// Extract OHLCV features
let price = current.price.unwrap_or(0.0) as f32;
let volume = current.quantity.unwrap_or(0.0) as f32;
let high = current.high.unwrap_or(price as f64) as f32;
let low = current.low.unwrap_or(price as f64) as f32;
// Calculate technical indicators
let bid_ask_spread = (high - low) / price.max(1e-6);
let momentum =
(price - prev.price.unwrap_or(price as f64) as f32) / price.max(1e-6);
// Calculate volatility from window
let prices: Vec<f32> = window.iter().map(|e| e.price.unwrap_or(0.0) as f32).collect();
let mean_price = prices.iter().sum::<f32>() / prices.len() as f32;
let variance = prices
.iter()
.map(|p| (p - mean_price).powi(2))
.sum::<f32>()
/ prices.len() as f32;
let volatility = variance.sqrt();
// Build state vector (pad or truncate to state_dim)
let mut state = vec![price, volume, bid_ask_spread, momentum, volatility];
state.resize(state_dim, 0.0);
state
})
.collect()
}
/// Convert market events to TFT/MAMBA-2 time series sequences
///
/// Returns (sequences, targets) where:
/// - sequences: [batch, seq_len, features]
/// - targets: [batch, prediction_horizon]
pub(crate) fn events_to_time_series(
events: &[MarketDataEvent],
seq_len: usize,
features_per_event: usize,
) -> Vec<Vec<f32>> {
if events.is_empty() || events.len() < seq_len {
return vec![];
}
events
.windows(seq_len)
.map(|window| {
window
.iter()
.flat_map(|event| {
let price = event.price.unwrap_or(0.0) as f32;
let volume = event.quantity.unwrap_or(0.0) as f32;
let high = event.high.unwrap_or(price as f64) as f32;
let low = event.low.unwrap_or(price as f64) as f32;
let open = event.open.unwrap_or(price as f64) as f32;
// Build feature vector [open, high, low, close, volume, ...]
let mut features = vec![open, high, low, price, volume];
features.resize(features_per_event, 0.0);
features
})
.collect()
})
.collect()
}
/// Load DQN states (convenience wrapper)
pub(crate) async fn load_dqn_states(count: usize, state_dim: usize) -> Result<Vec<Vec<f32>>> {
let loader = RealDataLoader::new();
// Check if files exist, otherwise return empty (tests will skip)
if !loader.files_exist() {
return Ok(vec![]);
}
let events = loader.load_btc_events(count + 5).await?; // +5 for window calculation
Ok(events_to_dqn_states(&events, state_dim))
}
/// Load TFT/MAMBA-2 sequences (convenience wrapper)
pub(crate) async fn load_tft_sequences(
count: usize,
seq_len: usize,
features_per_event: usize,
) -> Result<Vec<Vec<f32>>> {
let loader = RealDataLoader::new();
// Check if files exist, otherwise return empty (tests will skip)
if !loader.files_exist() {
return Ok(vec![]);
}
let events = loader.load_btc_events(count + seq_len).await?;
Ok(events_to_time_series(&events, seq_len, features_per_event))
}
/// Check if real data is available for tests
pub(crate) async fn real_data_available() -> bool {
let loader = RealDataLoader::new();
loader.files_exist()
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_real_data_loader_creation() {
let loader = RealDataLoader::new();
assert!(!loader.base_path.is_empty());
}
#[tokio::test]
#[ignore = "requires real Parquet files"]
async fn test_load_btc_events() {
let loader = RealDataLoader::new();
if !loader.files_exist() {
eprintln!("Skipping test: real data files not found");
return;
}
let events = loader.load_btc_events(100).await.unwrap();
assert_eq!(events.len(), 100);
assert!(events[0].price.is_some());
assert!(events[0].quantity.is_some());
}
#[tokio::test]
#[ignore = "requires real Parquet files"]
async fn test_events_to_dqn_states() {
let loader = RealDataLoader::new();
if !loader.files_exist() {
eprintln!("Skipping test: real data files not found");
return;
}
let events = loader.load_btc_events(50).await.unwrap();
let states = events_to_dqn_states(&events, 32);
assert!(!states.is_empty());
assert_eq!(states[0].len(), 32);
// Verify state contains non-zero values
assert!(states[0].iter().any(|&v| v != 0.0));
}
#[tokio::test]
#[ignore = "requires real Parquet files"]
async fn test_events_to_time_series() {
let loader = RealDataLoader::new();
if !loader.files_exist() {
eprintln!("Skipping test: real data files not found");
return;
}
let events = loader.load_btc_events(50).await.unwrap();
let sequences = events_to_time_series(&events, 10, 5);
assert!(!sequences.is_empty());
assert_eq!(sequences[0].len(), 10 * 5); // seq_len * features_per_event
// Verify sequences contain non-zero values
assert!(sequences[0].iter().any(|&v| v != 0.0));
}
#[tokio::test]
async fn test_load_dqn_states_wrapper() {
let states = load_dqn_states(50, 32).await.unwrap();
// If no data, should return empty (tests will skip)
if states.is_empty() {
eprintln!("No real data available, test will skip");
return;
}
assert_eq!(states[0].len(), 32);
}
#[tokio::test]
async fn test_load_tft_sequences_wrapper() {
let sequences = load_tft_sequences(30, 10, 5).await.unwrap();
// If no data, should return empty (tests will skip)
if sequences.is_empty() {
eprintln!("No real data available, test will skip");
return;
}
assert_eq!(sequences[0].len(), 10 * 5);
}
#[tokio::test]
async fn test_real_data_available() {
let available = real_data_available().await;
println!("Real data available: {}", available);
}
}