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>
688 lines
23 KiB
Rust
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);
|
|
}
|
|
}
|