Files
foxhunt/crates/ml/src/data_loaders/streaming_dbn_loader.rs
jgrusewski 09c515e3e9 fix(clippy): ZERO errors across entire workspace — CI ready
Final 31 ml crate fixes: unsafe_code allows, unused vars prefixed,
boolean simplification, dead code removal, integer suffix, drop cleanup.

cargo fix auto-removed ~30 unused imports from ml crate.

Total clippy cleanup: 278 errors → 0 across all ML crates.
Full workspace: `cargo clippy --workspace --lib -- -D warnings` = 0 errors.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-03-19 01:04:09 +01:00

688 lines
23 KiB
Rust

//! Streaming DBN Loader for Memory-Efficient Training
//!
//! Implements iterator-based streaming data loading to reduce memory usage from >2GB to <512MB.
//! Loads DBN files on-demand in batches of 10K bars and creates sequences incrementally.
//!
//! ## Key Features
//!
//! - **Memory Efficient**: <512MB peak memory (vs 2GB+ batch loading)
//! - **Iterator Pattern**: On-demand data loading with lazy evaluation
//! - **Sliding Window**: Incremental sequence creation from streaming data
//! - **Configurable Batching**: Adjustable batch size (default: 10,000 bars)
//! - **Maintains Performance**: <10% slower than batch loading
//!
//! ## Usage
//!
//! ```no_run
//! use ml::data_loaders::StreamingDbnLoader;
//! use NativeDevice;
//!
//! # async fn example() -> anyhow::Result<()> {
//! let loader = StreamingDbnLoader::new(60, 256).await?;
//! let stream = loader.stream_sequences("test_data/real/databento/ml_training", 0.9).await?;
//!
//! // Process sequences on-demand
//! for batch in stream {
//! let sequences = batch?;
//! println!("Processing {} sequences", sequences.len());
//! // Train model with batch...
//! }
//! # Ok(())
//! # }
//! ```
//!
//! ## Memory Comparison
//!
//! | Approach | Memory Usage | Speed |
//! |----------|--------------|-------|
//! | Batch | 2GB+ | 100% |
//! | Streaming| <512MB | ~95% |
use anyhow::{Context, Result};
use ml_core::native_types::NativeDevice;
use ml_core::cuda_autograd::GpuTensor;
use data::providers::databento::dbn_parser::{DbnParser, ProcessedMessage};
use dbn::decode::{DbnDecoder, DbnMetadata, DecodeRecordRef};
use rust_decimal::prelude::*;
use std::collections::{HashMap, VecDeque};
use std::path::{Path, PathBuf};
use tokio::fs;
use tracing::{info, warn};
/// Streaming DBN sequence loader for memory-efficient training
pub struct StreamingDbnLoader {
/// DBN parser for reading binary files
parser: DbnParser,
/// Target sequence length
seq_len: usize,
/// Feature dimension (d_model)
d_model: usize,
/// NativeDevice for tensor creation
device: NativeDevice,
/// Feature statistics for normalization
stats: FeatureStats,
/// Batch size for streaming (number of bars to load at once)
batch_size: usize,
/// Stride for sliding window (1 = every bar, 10 = every 10th bar)
stride: usize,
}
impl std::fmt::Debug for StreamingDbnLoader {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("StreamingDbnLoader")
.field("seq_len", &self.seq_len)
.field("d_model", &self.d_model)
.field("batch_size", &self.batch_size)
.field("stride", &self.stride)
.field("stats", &self.stats)
.finish_non_exhaustive()
}
}
/// Feature statistics for normalization
#[derive(Debug, Clone)]
struct FeatureStats {
price_mean: f64,
price_std: f64,
volume_mean: f64,
volume_std: f64,
}
impl Default for FeatureStats {
fn default() -> Self {
Self {
price_mean: 0.0,
price_std: 1.0,
volume_mean: 0.0,
volume_std: 1.0,
}
}
}
/// Streaming sequence iterator
pub struct SequenceStream {
/// DBN files to process
dbn_files: VecDeque<PathBuf>,
/// Current file being processed
current_file: Option<PathBuf>,
/// Message buffer for current file
message_buffer: VecDeque<ProcessedMessage>,
/// Loader reference
loader: StreamingDbnLoader,
/// Train/validation split ratio
train_split: f64,
/// Current position (for train/val split)
position: usize,
/// Total estimated sequences
total_sequences: usize,
/// Whether we're in training phase
is_training: bool,
}
impl std::fmt::Debug for SequenceStream {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SequenceStream")
.field("dbn_files_count", &self.dbn_files.len())
.field("current_file", &self.current_file)
.field("message_buffer_size", &self.message_buffer.len())
.field("loader", &self.loader)
.field("train_split", &self.train_split)
.field("position", &self.position)
.field("total_sequences", &self.total_sequences)
.field("is_training", &self.is_training)
.finish()
}
}
impl StreamingDbnLoader {
/// Create new streaming DBN sequence loader
///
/// # Arguments
/// * `seq_len` - Target sequence length (60-128 recommended)
/// * `d_model` - Feature dimension for MAMBA-2 (256, 512, or 1024)
///
/// # Returns
/// Configured streaming loader ready to process DBN files
pub async fn new(seq_len: usize, d_model: usize) -> Result<Self> {
let parser =
DbnParser::new().map_err(|e| anyhow::anyhow!("Failed to create DBN parser: {}", e))?;
// Configure symbol map for 6E.FUT (Euro FX futures)
let mut symbol_map = HashMap::new();
symbol_map.insert(0, "6E.FUT".to_owned());
symbol_map.insert(1, "6E.FUT".to_owned());
parser.update_symbol_map(symbol_map);
// Configure price scales (4 decimal places for FX)
let mut price_scales = HashMap::new();
price_scales.insert(0, 4);
price_scales.insert(1, 4);
parser.update_price_scales(price_scales);
let device = NativeDevice::Cuda(0);
// Default: 10,000 bars per batch (configurable)
// This provides good memory efficiency while maintaining performance
let batch_size = 10_000;
let stride = 100; // Sample every 100th bar
info!(
"Streaming DBN loader initialized (seq_len={}, d_model={}, batch_size={}, stride={}, device={:?})",
seq_len, d_model, batch_size, stride, device
);
Ok(Self {
parser,
seq_len,
d_model,
device,
stats: FeatureStats::default(),
batch_size,
stride,
})
}
/// Create loader with custom batch size and stride
pub async fn with_config(
seq_len: usize,
d_model: usize,
batch_size: usize,
stride: usize,
) -> Result<Self> {
let mut loader = Self::new(seq_len, d_model).await?;
loader.batch_size = batch_size.max(1000); // Minimum 1000 bars
loader.stride = stride.max(1); // Minimum stride of 1
info!(
"Custom config set: batch_size={}, stride={}",
loader.batch_size, loader.stride
);
Ok(loader)
}
/// Stream sequences from directory of DBN files
///
/// # Arguments
/// * `dbn_dir` - Directory containing .dbn files
/// * `train_split` - Fraction of data for training (0.0-1.0)
///
/// # Returns
/// Iterator over sequence batches (each batch contains multiple sequences)
pub async fn stream_sequences<P: AsRef<Path>>(
mut self,
dbn_dir: P,
train_split: f64,
) -> Result<SequenceStream> {
let path = dbn_dir.as_ref();
info!("📡 Streaming DBN sequences from: {:?}", path);
info!(
" Configuration: seq_len={}, d_model={}, batch_size={}, stride={}",
self.seq_len, self.d_model, self.batch_size, self.stride
);
// Find all .dbn files
let mut dbn_files = VecDeque::new();
let mut entries = fs::read_dir(path)
.await
.with_context(|| format!("Failed to read directory: {:?}", path))?;
while let Some(entry) = entries.next_entry().await? {
let path = entry.path();
if path.extension().and_then(|s| s.to_str()) == Some("dbn") {
dbn_files.push_back(path);
}
}
dbn_files.make_contiguous().sort();
info!("📁 Found {} DBN files for streaming", dbn_files.len());
if dbn_files.is_empty() {
return Err(anyhow::anyhow!("No DBN files found in {:?}", path));
}
// Estimate total sequences from all files (for progress tracking)
let total_sequences = self.estimate_total_sequences(&dbn_files).await?;
info!("📊 Estimated total sequences: ~{}", total_sequences);
// Compute feature statistics from a sample of files (first 10%)
info!("📊 Computing feature statistics from sample...");
self.compute_stats_from_sample(&dbn_files).await?;
info!(
" price_mean={:.2}, price_std={:.2}, volume_mean={:.2}, volume_std={:.2}",
self.stats.price_mean,
self.stats.price_std,
self.stats.volume_mean,
self.stats.volume_std
);
Ok(SequenceStream {
dbn_files,
current_file: None,
message_buffer: VecDeque::new(),
loader: self,
train_split,
position: 0,
total_sequences,
is_training: true,
})
}
/// Estimate total sequences from all files (quick scan)
async fn estimate_total_sequences(&self, files: &VecDeque<PathBuf>) -> Result<usize> {
let mut total_messages = 0;
// Sample first 10 files to estimate
let sample_size = files.len().min(10);
for file_path in files.iter().take(sample_size) {
let messages = self.count_messages_in_file(file_path).await?;
total_messages += messages;
}
// Extrapolate to all files
let avg_messages_per_file = total_messages / sample_size.max(1);
let estimated_total = avg_messages_per_file * files.len();
// Estimate sequences with stride
let estimated_sequences = estimated_total / self.stride;
Ok(estimated_sequences)
}
/// Count messages in a DBN file (quick scan without full parsing)
async fn count_messages_in_file(&self, path: &Path) -> Result<usize> {
use std::fs::File;
use std::io::BufReader;
let file = File::open(path).with_context(|| format!("Failed to open: {:?}", path))?;
let reader = BufReader::new(file);
let mut decoder = DbnDecoder::new(reader)
.map_err(|e| anyhow::anyhow!("Failed to create DBN decoder: {}", e))?;
let mut count = 0;
loop {
match decoder.decode_record_ref() {
Ok(Some(_)) => count += 1,
Ok(None) => break,
Err(_) => break,
}
}
Ok(count)
}
/// Compute feature statistics from a sample of files
async fn compute_stats_from_sample(&mut self, files: &VecDeque<PathBuf>) -> Result<()> {
let mut prices = Vec::new();
let mut volumes = Vec::new();
// Sample first 10% of files
let sample_size = (files.len() / 10).max(1);
for file_path in files.iter().take(sample_size) {
let messages = self.load_file(file_path).await?;
for msg in messages {
match msg {
ProcessedMessage::Ohlcv { close, volume, .. } => {
prices.push(close.to_f64());
volumes.push(volume.to_f64().unwrap_or(0.0));
},
ProcessedMessage::Trade { price, size, .. } => {
prices.push(price.to_f64());
volumes.push(size.to_f64().unwrap_or(0.0));
},
ProcessedMessage::Quote { .. }
| ProcessedMessage::OrderBook { .. }
| ProcessedMessage::Status { .. } => {},
}
}
}
if prices.is_empty() {
return Err(anyhow::anyhow!("No price data found for normalization"));
}
// Compute mean and std
let price_mean = prices.iter().sum::<f64>() / prices.len() as f64;
let price_var =
prices.iter().map(|p| (p - price_mean).powi(2)).sum::<f64>() / prices.len() as f64;
let price_std = price_var.sqrt().max(1e-8);
let volume_mean = volumes.iter().sum::<f64>() / volumes.len() as f64;
let volume_var = volumes
.iter()
.map(|v| (v - volume_mean).powi(2))
.sum::<f64>()
/ volumes.len() as f64;
let volume_std = volume_var.sqrt().max(1e-8);
self.stats = FeatureStats {
price_mean,
price_std,
volume_mean,
volume_std,
};
Ok(())
}
/// Load messages from a single DBN file
async fn load_file(&self, path: &Path) -> Result<Vec<ProcessedMessage>> {
use std::fs::File;
use std::io::BufReader;
let file = File::open(path).with_context(|| format!("Failed to open: {:?}", path))?;
let reader = BufReader::new(file);
let mut decoder = DbnDecoder::new(reader)
.map_err(|e| anyhow::anyhow!("Failed to create DBN decoder: {}", e))?;
let metadata = decoder.metadata();
let symbol = metadata
.symbols
.first()
.cloned()
.unwrap_or_else(|| "UNKNOWN".to_owned());
let mut messages = Vec::new();
loop {
match decoder.decode_record_ref() {
Ok(Some(record)) => {
let record_enum = record
.as_enum()
.map_err(|e| anyhow::anyhow!("Failed to convert record: {}", e))?;
match record_enum {
dbn::RecordRefEnum::Ohlcv(ohlcv) => {
let open_f64 = ohlcv.open as f64 * 1e-9;
let high_f64 = ohlcv.high as f64 * 1e-9;
let low_f64 = ohlcv.low as f64 * 1e-9;
let close_f64 = ohlcv.close as f64 * 1e-9;
let open = common::Price::from_f64(open_f64.abs())?;
let high = common::Price::from_f64(high_f64.abs())?;
let low = common::Price::from_f64(low_f64.abs())?;
let close = common::Price::from_f64(close_f64.abs())?;
let volume = Decimal::from(ohlcv.volume);
use data::HardwareTimestamp;
let timestamp = HardwareTimestamp::from_nanos(ohlcv.hd.ts_event);
messages.push(ProcessedMessage::Ohlcv {
symbol: symbol.clone(),
open,
high,
low,
close,
volume,
timestamp,
});
},
dbn::RecordRefEnum::Mbo(_)
| dbn::RecordRefEnum::Trade(_)
| dbn::RecordRefEnum::Mbp1(_)
| dbn::RecordRefEnum::Mbp10(_)
| dbn::RecordRefEnum::Status(_)
| dbn::RecordRefEnum::InstrumentDef(_)
| dbn::RecordRefEnum::Imbalance(_)
| dbn::RecordRefEnum::Stat(_)
| dbn::RecordRefEnum::Error(_)
| dbn::RecordRefEnum::SymbolMapping(_)
| dbn::RecordRefEnum::System(_)
| dbn::RecordRefEnum::Cmbp1(_)
| dbn::RecordRefEnum::Bbo(_)
| dbn::RecordRefEnum::Cbbo(_) => {}, // Skip other message types
}
},
Ok(None) => break,
Err(e) => {
warn!("Error decoding record: {}", e);
break;
},
}
}
Ok(messages)
}
/// Extract normalized features from a message
fn extract_features(&self, msg: &ProcessedMessage) -> Result<Vec<f32>> {
match msg {
ProcessedMessage::Ohlcv {
open,
high,
low,
close,
volume,
..
} => {
let o = (open.to_f64() - self.stats.price_mean) / self.stats.price_std;
let h = (high.to_f64() - self.stats.price_mean) / self.stats.price_std;
let l = (low.to_f64() - self.stats.price_mean) / self.stats.price_std;
let c = (close.to_f64() - self.stats.price_mean) / self.stats.price_std;
let v = (volume.to_f64().unwrap_or(0.0) - self.stats.volume_mean)
/ self.stats.volume_std;
// Derived features
let range = h - l;
let body = c - o;
let upper_wick = h - c.max(o);
let lower_wick = l.min(o) - l;
Ok(vec![
o as f32,
h as f32,
l as f32,
c as f32,
v as f32,
range as f32,
body as f32,
upper_wick as f32,
lower_wick as f32,
])
},
ProcessedMessage::Trade { .. }
| ProcessedMessage::Quote { .. }
| ProcessedMessage::OrderBook { .. }
| ProcessedMessage::Status { .. } => Ok(vec![0.0; 9]),
}
}
/// Create sequence from message window
fn create_sequence(&self, window: &[ProcessedMessage]) -> Result<(GpuTensor, GpuTensor)> {
if window.len() != self.seq_len + 1 {
return Err(anyhow::anyhow!(
"Invalid window size: {} (expected {})",
window.len(),
self.seq_len + 1
));
}
// Extract features for seq_len steps
let mut features = Vec::with_capacity(self.seq_len * self.d_model);
for msg in &window[..self.seq_len] {
let msg_features = self.extract_features(msg)?;
// Pad or truncate to d_model dimension
for j in 0..self.d_model {
if j < msg_features.len() {
features.push(msg_features[j]);
} else {
features.push(0.0);
}
}
}
// Target is next timestep
let target_msg = &window[self.seq_len];
let target_features = self.extract_features(target_msg)?;
let mut target = vec![0.0; self.d_model];
let copy_len = self.d_model.min(target_features.len());
target[..copy_len].copy_from_slice(&target_features[..copy_len]);
// Create tensors with batch dimension [batch=1, seq_len, d_model]
let ordinal = self.device.cuda_ordinal().unwrap_or(0);
let ml_dev = ml_core::device::MlDevice::cuda(ordinal)
.map_err(|e| anyhow::anyhow!("CUDA device for streaming loader: {e}"))?;
let stream = ml_dev.cuda_stream()
.map_err(|e| anyhow::anyhow!("CUDA stream for streaming loader: {e}"))?;
let input = GpuTensor::from_host(&features, vec![1, self.seq_len, self.d_model], stream)?;
let target_tensor = GpuTensor::from_host(&target, vec![1, 1, self.d_model], stream)?;
Ok((input, target_tensor))
}
}
impl SequenceStream {
/// Load next file into message buffer
async fn load_next_file(&mut self) -> Result<bool> {
if self.dbn_files.is_empty() {
return Ok(false);
}
let file_path = match self.dbn_files.pop_front() {
Some(path) => path,
None => return Ok(false),
};
info!(
"📖 Loading file: {:?}",
file_path.file_name().unwrap_or_default()
);
let messages = self.loader.load_file(&file_path).await?;
info!(" Loaded {} messages", messages.len());
self.message_buffer.extend(messages);
self.current_file = Some(file_path);
Ok(true)
}
/// Get next batch of sequences
pub async fn next_batch(&mut self) -> Result<Option<Vec<(GpuTensor, GpuTensor)>>> {
let mut sequences = Vec::new();
let target_batch_size = self.loader.batch_size / self.loader.stride;
// Keep creating sequences until we have enough or run out of data
while sequences.len() < target_batch_size {
// Ensure we have enough messages in buffer
while self.message_buffer.len() < self.loader.seq_len + 1 {
if !self.load_next_file().await? {
// No more files - return what we have
if sequences.is_empty() {
return Ok(None);
} else {
return Ok(Some(sequences));
}
}
}
// Create sequence from sliding window
let window: Vec<_> = self
.message_buffer
.iter()
.take(self.loader.seq_len + 1)
.cloned()
.collect();
if window.len() == self.loader.seq_len + 1 {
match self.loader.create_sequence(&window) {
Ok(seq) => {
sequences.push(seq);
self.position += 1;
},
Err(e) => {
warn!("Failed to create sequence: {}", e);
},
}
}
// Advance window by stride
for _ in 0..self.loader.stride.min(self.message_buffer.len()) {
self.message_buffer.pop_front();
}
// Check if we've exhausted buffer
if self.message_buffer.len() < self.loader.seq_len + 1 {
if self.dbn_files.is_empty() {
// No more data
break;
}
}
}
if sequences.is_empty() {
Ok(None)
} else {
info!(
"✅ Created batch of {} sequences (position: {})",
sequences.len(),
self.position
);
Ok(Some(sequences))
}
}
/// Reset to validation phase (for train/val split)
pub async fn switch_to_validation(&mut self) -> Result<()> {
self.is_training = false;
self.position = 0;
self.message_buffer.clear();
// Skip to validation split
let train_sequences = (self.total_sequences as f64 * self.train_split) as usize;
info!(
"📊 Switching to validation (skipping {} training sequences)",
train_sequences
);
// Fast-forward to validation data
while self.position < train_sequences {
if self.next_batch().await?.is_none() {
break;
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_streaming_loader_creation() {
let loader = StreamingDbnLoader::new(60, 256).await;
assert!(loader.is_ok());
}
#[tokio::test]
async fn test_custom_config() {
let loader = StreamingDbnLoader::with_config(60, 256, 5000, 50).await;
assert!(loader.is_ok());
let loader = loader.unwrap();
assert_eq!(loader.batch_size, 5000);
assert_eq!(loader.stride, 50);
}
}