//! DBN Sequence Loader for MAMBA-2 Training //! //! Loads real market data from Databento DBN files and creates sequences for MAMBA-2 //! state space model training. Handles OHLCV data with microstructure features. //! //! ## Features //! //! - Loads DBN files via DbnParser //! - Creates fixed-length sequences (60-128 timesteps) //! - Extracts price, volume, spreads, and order flow features //! - Maintains temporal ordering for state space models //! - Normalizes features for neural network input //! //! ## Usage //! //! ```no_run //! use ml::data_loaders::DbnSequenceLoader; //! use candle_core::Device; //! //! # async fn example() -> anyhow::Result<()> { //! let loader = DbnSequenceLoader::new(60, 256).await?; //! let (train_data, val_data) = loader //! .load_sequences("test_data/real/databento/ml_training_small", 0.9) //! .await?; //! //! println!("Loaded {} training sequences", train_data.len()); //! # Ok(()) //! # } //! ``` use anyhow::{Context, Result}; use candle_core::{DType, Device, Tensor}; use data::providers::databento::dbn_parser::{DbnParser, ProcessedMessage}; use dbn::decode::{DbnDecoder, DbnMetadata}; use rust_decimal::prelude::*; use std::collections::HashMap; use std::path::Path; use tokio::fs; use tracing::{debug, info, warn}; use crate::data_loaders::dbn_tick_adapter::Tick; use crate::features::alternative_bars::{ DollarBarSampler, ImbalanceBarSampler, RunBarSampler, TickBarSampler, VolumeBarSampler, }; use crate::features::normalization::FeatureNormalizer; use crate::types::OHLCVBar; // Wave D regime detection feature extractors use crate::ensemble::MarketRegime; use crate::features::extraction::extract_ml_features; use crate::features::regime_adaptive::RegimeAdaptiveFeatures; use crate::features::regime_adx::RegimeADXFeatures; use crate::features::regime_cusum::RegimeCUSUMFeatures; use crate::features::regime_transition::RegimeTransitionFeatures; /// Bar sampling method for alternative bar types (Wave B) #[derive(Debug, Clone)] pub enum BarSamplingMethod { /// Time-based bars (default) - fixed time intervals TimeBars, /// Tick bars - fixed number of ticks TickBars(usize), /// Volume bars - fixed volume threshold VolumeBars(f64), /// Dollar bars - fixed dollar value threshold DollarBars(f64), /// Imbalance bars - fixed imbalance threshold ImbalanceBars(f64), /// Run bars - consecutive directional ticks RunBars(usize), } /// DBN sequence loader for MAMBA-2 training pub struct DbnSequenceLoader { /// DBN parser for reading binary files parser: DbnParser, /// Target sequence length seq_len: usize, /// Feature dimension (d_model) - dynamically computed from feature_config d_model: usize, /// Device for tensor creation device: Device, /// Feature statistics for normalization stats: FeatureStats, /// Maximum sequences per symbol (prevents memory overflow) max_sequences_per_symbol: Option, /// Stride for sliding window (1 = every bar, 10 = every 10th bar) stride: usize, /// Feature configuration (Wave A/B/C) feature_config: crate::features::config::FeatureConfig, /// Bar sampling method (Wave B alternative bars) bar_sampling_method: BarSamplingMethod, /// Feature normalizer (Wave C/D normalization) normalizer: FeatureNormalizer, /// Wave D regime detection feature extractors (24 features: indices 201-224) regime_cusum: RegimeCUSUMFeatures, // 10 features (201-210) regime_adx: RegimeADXFeatures, // 5 features (211-215) regime_transition: RegimeTransitionFeatures, // 5 features (216-220) regime_adaptive: RegimeAdaptiveFeatures, // 4 features (221-224) /// Current detected regime for transition tracking current_regime: MarketRegime, /// OHLCV bar buffer for ADX (requires historical bars with i64 timestamp) bar_buffer_adx: Vec, /// OHLCV bar buffer for adaptive features (requires historical bars with DateTime timestamp) bar_buffer_adaptive: Vec, } impl std::fmt::Debug for DbnSequenceLoader { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("DbnSequenceLoader") .field("seq_len", &self.seq_len) .field("d_model", &self.d_model) .field("max_sequences_per_symbol", &self.max_sequences_per_symbol) .field("stride", &self.stride) .field("stats", &self.stats) .field("feature_config", &self.feature_config) .finish_non_exhaustive() } } /// Feature statistics for normalization #[derive(Debug, Clone)] struct FeatureStats { /// Price mean price_mean: f64, /// Price standard deviation price_std: f64, /// Volume mean volume_mean: f64, /// Volume standard deviation 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, } } } impl DbnSequenceLoader { /// Create new DBN sequence loader with default Wave A config (26 features) /// /// # Arguments /// * `seq_len` - Target sequence length (60-128 recommended) /// * `d_model` - Feature dimension for MAMBA-2 (must match feature_config.feature_count()) /// /// # Returns /// Configured loader ready to process DBN files /// /// # Note /// This constructor uses Wave A config (26 features). For other configs, use `with_feature_config()`. pub async fn new(seq_len: usize, d_model: usize) -> Result { use std::collections::HashMap; let feature_config = crate::features::config::FeatureConfig::wave_a(); // Validate d_model matches feature_config if d_model != feature_config.feature_count() { anyhow::bail!( "d_model ({}) does not match feature_config.feature_count() ({}). Use wave_a()={}, wave_b()={}, wave_c()={}+", d_model, feature_config.feature_count(), crate::features::config::FeatureConfig::wave_a().feature_count(), crate::features::config::FeatureConfig::wave_b().feature_count(), crate::features::config::FeatureConfig::wave_c().feature_count() ); } 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_string()); symbol_map.insert(1, "6E.FUT".to_string()); 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 = Device::cuda_if_available(0).unwrap_or(Device::Cpu); // W12-15B: Increased to 10,000 sequences per symbol (10x increase with stride=10) // For 204K bars (6E.FUT) with seq_len=60 and stride=10: ~20,400 possible sequences // Limited to 10K for memory efficiency (~512MB per symbol) let max_sequences_per_symbol = Some(10_000); let stride = 10; // Sample every 10th bar for better training data coverage (W12-15B: 10x more sequences) info!("DBN sequence loader initialized (seq_len={}, d_model={}, feature_phase={:?}, device={:?}, max_sequences={:?}, stride={})", seq_len, d_model, feature_config.phase, device, max_sequences_per_symbol, stride); // Initialize feature normalizer with custom window sizes for Wave D let normalizer = FeatureNormalizer::with_config( 50, // price_window 50, // volume_window 20, // microstructure_window 30, // regime_window (Wave D) ); // Initialize Wave D regime detection extractors // CUSUM: target_mean=0.0 (log returns), target_std=1.0, drift=0.5, threshold=4.0 let regime_cusum = RegimeCUSUMFeatures::new(0.0, 1.0, 0.5, 4.0); // ADX: period=14 (standard) let regime_adx = RegimeADXFeatures::new(14); // Transition matrix: 4 regimes, alpha=0.1 let regime_transition = RegimeTransitionFeatures::new(4, 0.1); // Adaptive features: window=20, max_position=100K, atr_period=14 let regime_adaptive = RegimeAdaptiveFeatures::new(20, 100_000.0, 14); Ok(Self { parser, seq_len, d_model, device, stats: FeatureStats::default(), max_sequences_per_symbol, stride, feature_config, bar_sampling_method: BarSamplingMethod::TimeBars, normalizer, regime_cusum, regime_adx, regime_transition, regime_adaptive, current_regime: MarketRegime::Normal, bar_buffer_adx: Vec::with_capacity(100), bar_buffer_adaptive: Vec::with_capacity(100), }) } /// Create new DBN sequence loader with custom feature configuration /// /// # Arguments /// * `seq_len` - Target sequence length (60-128 recommended) /// * `feature_config` - Feature configuration (Wave A/B/C) /// /// # Returns /// Configured loader with specified feature configuration pub async fn with_feature_config( seq_len: usize, feature_config: crate::features::config::FeatureConfig, ) -> Result { use std::collections::HashMap; let d_model = feature_config.feature_count(); 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_string()); symbol_map.insert(1, "6E.FUT".to_string()); 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 = Device::cuda_if_available(0).unwrap_or(Device::Cpu); // W12-15B: Increased to 10,000 sequences per symbol (10x increase with stride=10) // For 204K bars (6E.FUT) with seq_len=60 and stride=10: ~20,400 possible sequences // Limited to 10K for memory efficiency (~512MB per symbol) let max_sequences_per_symbol = Some(10_000); let stride = 10; // Sample every 10th bar for better training data coverage (W12-15B: 10x more sequences) info!("DBN sequence loader initialized (seq_len={}, d_model={}, feature_phase={:?}, device={:?}, max_sequences={:?}, stride={})", seq_len, d_model, feature_config.phase, device, max_sequences_per_symbol, stride); // Initialize feature normalizer with custom window sizes for Wave D let normalizer = FeatureNormalizer::with_config( 50, // price_window 50, // volume_window 20, // microstructure_window 30, // regime_window (Wave D) ); // Initialize Wave D regime detection extractors // CUSUM: target_mean=0.0 (log returns), target_std=1.0, drift=0.5, threshold=4.0 let regime_cusum = RegimeCUSUMFeatures::new(0.0, 1.0, 0.5, 4.0); // ADX: period=14 (standard) let regime_adx = RegimeADXFeatures::new(14); // Transition matrix: 4 regimes, alpha=0.1 let regime_transition = RegimeTransitionFeatures::new(4, 0.1); // Adaptive features: window=20, max_position=100K, atr_period=14 let regime_adaptive = RegimeAdaptiveFeatures::new(20, 100_000.0, 14); Ok(Self { parser, seq_len, d_model, device, stats: FeatureStats::default(), max_sequences_per_symbol, stride, feature_config, bar_sampling_method: BarSamplingMethod::TimeBars, normalizer, regime_cusum, regime_adx, regime_transition, regime_adaptive, current_regime: MarketRegime::Normal, bar_buffer_adx: Vec::with_capacity(100), bar_buffer_adaptive: Vec::with_capacity(100), }) } /// Set bar sampling method (Wave B alternative bars) /// /// # Arguments /// * `method` - Bar sampling method (TimeBars, TickBars, VolumeBars, etc.) /// /// # Example /// ```no_run /// # use ml::data_loaders::{DbnSequenceLoader, BarSamplingMethod}; /// # async fn example() -> anyhow::Result<()> { /// let mut loader = DbnSequenceLoader::new(60, 26).await?; /// loader.set_bar_sampling_method(BarSamplingMethod::DollarBars(2_000_000.0)); /// # Ok(()) /// # } /// ``` pub fn set_bar_sampling_method(&mut self, method: BarSamplingMethod) { info!("Setting bar sampling method: {:?}", method); self.bar_sampling_method = method; } /// Get current bar sampling method pub fn bar_sampling_method(&self) -> &BarSamplingMethod { &self.bar_sampling_method } /// Create new DBN sequence loader with custom limits /// /// # Arguments /// * `seq_len` - Target sequence length /// * `d_model` - Feature dimension (must match wave_a()=26, wave_b()=36, wave_c()=65+) /// * `max_sequences_per_symbol` - Maximum sequences per symbol (None = unlimited) /// * `stride` - Stride for sliding window (1 = every bar, 10 = every 10th bar) /// /// # Note /// This constructor uses Wave A config (26 features). For other configs, use `with_feature_config()`. pub async fn with_limits( seq_len: usize, d_model: usize, max_sequences_per_symbol: Option, stride: usize, ) -> Result { let mut loader = Self::new(seq_len, d_model).await?; loader.max_sequences_per_symbol = max_sequences_per_symbol; loader.stride = stride.max(1); // Ensure stride >= 1 info!( "Custom limits set: max_sequences={:?}, stride={}", max_sequences_per_symbol, loader.stride ); Ok(loader) } /// Load sequences from a directory of DBN files /// /// # Arguments /// * `dbn_dir` - Directory containing .dbn files /// * `train_split` - Fraction of data for training (0.0-1.0) /// /// # Returns /// Tuple of (train_sequences, val_sequences) as (input, target) pairs /// /// # Wave B Alternative Bars /// If `bar_sampling_method` is set to anything other than `TimeBars`, this method will: /// 1. Load DBN OHLCV bars /// 2. Convert to tick data using DBNTickAdapter /// 3. Apply alternative bar sampler (Tick/Volume/Dollar/Imbalance/Run) /// 4. Create sequences from alternative bars pub async fn load_sequences>( &mut self, dbn_dir: P, train_split: f64, ) -> Result<(Vec<(Tensor, Tensor)>, Vec<(Tensor, Tensor)>)> { let path = dbn_dir.as_ref(); info!("🔄 Loading DBN sequences from: {:?}", path); info!(" Configuration: seq_len={}, d_model={}, stride={}, max_sequences={:?}, bar_method={:?}", self.seq_len, self.d_model, self.stride, self.max_sequences_per_symbol, self.bar_sampling_method); // Find all .dbn files let mut dbn_files = Vec::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(path); } } dbn_files.sort(); info!("📁 Found {} DBN files", dbn_files.len()); if dbn_files.is_empty() { return Err(anyhow::anyhow!("No DBN files found in {:?}", path)); } // Load all messages from all files and group by symbol let mut symbol_messages: HashMap> = HashMap::new(); let total_files = dbn_files.len(); for (idx, file_path) in dbn_files.iter().enumerate() { let progress = ((idx + 1) as f64 / total_files as f64 * 100.0) as usize; info!( "📖 Processing file {}/{} ({}%): {:?}", idx + 1, total_files, progress, file_path.file_name().unwrap_or_default() ); let messages = self.load_file(file_path).await?; info!(" Loaded {} messages", messages.len()); // Group by symbol for temporal ordering for msg in messages { let symbol = self.get_message_symbol(&msg); symbol_messages.entry(symbol).or_default().push(msg); } // Log memory status every 50 files if (idx + 1) % 50 == 0 { let total_messages: usize = symbol_messages.values().map(|v| v.len()).sum(); info!( " 💾 Memory checkpoint: {} messages across {} symbols", total_messages, symbol_messages.len() ); } } let total_messages: usize = symbol_messages.values().map(|v| v.len()).sum(); info!( "✅ Loaded {} messages for {} symbols", total_messages, symbol_messages.len() ); // Apply alternative bar sampling if configured (Wave B) let symbol_messages = match &self.bar_sampling_method { BarSamplingMethod::TimeBars => { info!("📊 Using time-based bars (default)"); symbol_messages }, _ => { info!( "🔄 Applying alternative bar sampling: {:?}", self.bar_sampling_method ); self.apply_alternative_bar_sampling(symbol_messages).await? }, }; // Compute feature statistics for normalization info!("📊 Computing feature statistics..."); self.compute_stats(&symbol_messages)?; 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 ); // Create sequences from each symbol's messages info!("🔨 Creating sequences with sliding window..."); let mut all_sequences = Vec::new(); let total_symbols = symbol_messages.len(); for (sym_idx, (symbol, messages)) in symbol_messages.into_iter().enumerate() { let progress = ((sym_idx + 1) as f64 / total_symbols as f64 * 100.0) as usize; info!( " Processing symbol {}/{} ({}%): {} ({} messages)", sym_idx + 1, total_symbols, progress, symbol, messages.len() ); if messages.len() < self.seq_len + 1 { warn!( " ⚠️ Skipping {}: only {} messages (need {})", symbol, messages.len(), self.seq_len + 1 ); continue; } let sequences = self.create_sequences(&messages)?; info!(" Created {} sequences from {}", sequences.len(), symbol); all_sequences.extend(sequences); info!(" Total sequences so far: {}", all_sequences.len()); // Estimate memory usage let seq_memory_mb = (all_sequences.len() * self.seq_len * self.d_model * 4) / (1024 * 1024); info!( " 💾 Estimated memory: ~{}MB for {} sequences", seq_memory_mb, all_sequences.len() ); } info!("✅ Created {} total sequences", all_sequences.len()); if all_sequences.is_empty() { return Err(anyhow::anyhow!( "No sequences created! Check seq_len and data size." )); } // Split into train/val info!( "✂️ Splitting data (train={:.0}%, val={:.0}%)", train_split * 100.0, (1.0 - train_split) * 100.0 ); let split_idx = (all_sequences.len() as f64 * train_split) as usize; let train_data = all_sequences[..split_idx].to_vec(); let val_data = all_sequences[split_idx..].to_vec(); info!( "✅ Split complete: {} training, {} validation", train_data.len(), val_data.len() ); Ok((train_data, val_data)) } /// Load messages from a single DBN file using official dbn crate decoder /// /// This replaces the custom find_data_start() heuristic that only found 2 messages. /// Now properly decodes all OHLCV bars (400-500+ records per file). async fn load_file>(&self, path: P) -> Result> { use dbn::decode::DecodeRecordRef; use std::fs::File; use std::io::BufReader; let path = path.as_ref(); // Open file and create official DBN decoder 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))?; // Read metadata (for symbol mapping) let metadata = decoder.metadata(); let symbol = metadata .symbols .first() .map(|s| s.to_string()) .unwrap_or_else(|| "UNKNOWN".to_string()); debug!( "DBN file metadata: dataset={:?}, schema={:?}, symbol={}", metadata.dataset, metadata.schema, symbol ); // Decode all records and convert to ProcessedMessage let mut messages = Vec::new(); let mut ohlcv_count = 0; let mut other_count = 0; let mut idx = 0; loop { match decoder.decode_record_ref() { Ok(Some(record)) => { idx += 1; // Convert RecordRef to RecordRefEnum for pattern matching let record_enum = record .as_enum() .map_err(|e| anyhow::anyhow!("Failed to convert record to enum: {}", e))?; match record_enum { dbn::RecordRefEnum::Ohlcv(ohlcv) => { ohlcv_count += 1; // Prices are i64 scaled by 1e-9 per DBN specification 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; // Log first few records for validation if ohlcv_count <= 5 { debug!( "Raw OHLCV #{}: open={}, high={}, low={}, close={}", ohlcv_count, ohlcv.open, ohlcv.high, ohlcv.low, ohlcv.close ); debug!( "Scaled OHLCV #{}: open={:.6}, high={:.6}, low={:.6}, close={:.6}", ohlcv_count, open_f64, high_f64, low_f64, close_f64 ); } // Use absolute values for Price type (futures data can have negative values) 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); // Create HardwareTimestamp from ts_event (nanoseconds since Unix epoch) use trading_engine::timing::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::Trade(trade) => { other_count += 1; // Prices are i64 scaled by 1e-9 per DBN specification // Use absolute value for Price type (futures data can have negative values) let price_f64 = trade.price as f64 * 1e-9; let price = common::Price::from_f64(price_f64.abs())?; let size = Decimal::from(trade.size); use trading_engine::timing::HardwareTimestamp; let timestamp = HardwareTimestamp::from_nanos(trade.hd.ts_event); // Determine side from trade action/flags (c_char is i8) use common::OrderSide; let side = if trade.side == b'B' as i8 { OrderSide::Buy } else if trade.side == b'A' as i8 { OrderSide::Sell } else { OrderSide::Buy // Default }; messages.push(ProcessedMessage::Trade { symbol: symbol.clone(), price, size, side, trade_id: None, conditions: vec![], timestamp, }); }, dbn::RecordRefEnum::Mbp1(mbp) => { other_count += 1; // Mbp1Msg has price/size/side, not separate bid/ask fields // Prices are i64 scaled by 1e-9 per DBN specification // Use absolute value for Price type (futures data can have negative values) let price_f64 = mbp.price as f64 * 1e-9; let price = common::Price::from_f64(price_f64.abs())?; let size = Decimal::from(mbp.size); use trading_engine::timing::HardwareTimestamp; let timestamp = HardwareTimestamp::from_nanos(mbp.hd.ts_event); // Determine bid/ask from side field (c_char is i8) let (bid, ask, bid_size, ask_size) = if mbp.side == b'B' as i8 { // Bid side (Some(price), None, Some(size), None) } else if mbp.side == b'A' as i8 { // Ask side (None, Some(price), None, Some(size)) } else { // Unknown side - treat as bid (Some(price), None, Some(size), None) }; messages.push(ProcessedMessage::Quote { symbol: symbol.clone(), bid, ask, bid_size, ask_size, exchange: Some("UNKNOWN".to_string()), timestamp, }); }, _ => { // Skip other message types }, } }, Ok(None) => { // End of stream break; }, Err(e) => { return Err(anyhow::anyhow!("Failed to decode record {}: {}", idx, e)); }, } } info!( "Loaded {} OHLCV messages from {:?} ({} other messages)", ohlcv_count, path.file_name().unwrap_or_default(), other_count ); Ok(messages) } /// Apply alternative bar sampling to OHLCV messages (Wave B) /// /// Converts time-based OHLCV bars to alternative bar types: /// - TickBars: Fixed number of ticks /// - VolumeBars: Fixed volume threshold /// - DollarBars: Fixed dollar value threshold /// - ImbalanceBars: Fixed imbalance threshold /// - RunBars: Consecutive directional ticks /// /// # Process /// 1. Convert OHLCV bars to ticks (4 ticks per bar: OHLC) /// 2. Apply alternative bar sampler /// 3. Convert alternative bars back to ProcessedMessage::Ohlcv async fn apply_alternative_bar_sampling( &self, symbol_messages: HashMap>, ) -> Result>> { let mut result = HashMap::new(); for (symbol, messages) in symbol_messages.into_iter() { info!(" Processing {} messages for {}", messages.len(), symbol); // Convert OHLCV messages to ticks let ticks = self.ohlcv_messages_to_ticks(&messages)?; info!(" Converted to {} ticks", ticks.len()); // Apply alternative bar sampler let alternative_bars = self.apply_bar_sampler(&ticks)?; info!(" Generated {} alternative bars", alternative_bars.len()); // Convert alternative bars back to ProcessedMessage::Ohlcv let alternative_messages: Vec = alternative_bars .into_iter() .map(|bar| ProcessedMessage::Ohlcv { symbol: symbol.clone(), open: common::Price::from_f64(bar.open) .unwrap_or_else(|_| common::Price::from_f64(0.0).unwrap()), high: common::Price::from_f64(bar.high) .unwrap_or_else(|_| common::Price::from_f64(0.0).unwrap()), low: common::Price::from_f64(bar.low) .unwrap_or_else(|_| common::Price::from_f64(0.0).unwrap()), close: common::Price::from_f64(bar.close) .unwrap_or_else(|_| common::Price::from_f64(0.0).unwrap()), volume: Decimal::from_f64(bar.volume).unwrap_or(Decimal::ZERO), timestamp: trading_engine::timing::HardwareTimestamp::from_nanos( bar.timestamp.timestamp_nanos_opt().unwrap_or(0) as u64, ), }) .collect(); result.insert(symbol, alternative_messages); } Ok(result) } /// Convert OHLCV messages to ticks (4 ticks per bar: open, high, low, close) fn ohlcv_messages_to_ticks(&self, messages: &[ProcessedMessage]) -> Result> { let mut ticks = Vec::new(); for msg in messages { match msg { ProcessedMessage::Ohlcv { open, high, low, close, volume, timestamp, .. } => { // Convert timestamp to DateTime let ts_nanos = timestamp.as_nanos() as i64; let datetime = chrono::DateTime::from_timestamp_nanos(ts_nanos); // Distribute volume evenly across 4 ticks let volume_per_tick = volume.to_f64().unwrap_or(0.0) / 4.0; // Create 4 ticks per bar (OHLC) ticks.push(Tick { price: open.to_f64(), volume: volume_per_tick, timestamp: datetime, }); ticks.push(Tick { price: high.to_f64(), volume: volume_per_tick, timestamp: datetime, }); ticks.push(Tick { price: low.to_f64(), volume: volume_per_tick, timestamp: datetime, }); ticks.push(Tick { price: close.to_f64(), volume: volume_per_tick, timestamp: datetime, }); }, _ => { // Skip non-OHLCV messages }, } } Ok(ticks) } /// Apply bar sampler based on configured method fn apply_bar_sampler(&self, ticks: &[Tick]) -> Result> { let mut bars = Vec::new(); match &self.bar_sampling_method { BarSamplingMethod::TimeBars => { // This should never happen (already handled in load_sequences) anyhow::bail!("TimeBars should not reach apply_bar_sampler"); }, BarSamplingMethod::TickBars(threshold) => { let mut sampler = TickBarSampler::new(*threshold); for tick in ticks { if let Some(bar) = sampler.update(tick.price, tick.volume, tick.timestamp) { bars.push(bar); } } }, BarSamplingMethod::VolumeBars(threshold) => { let mut sampler = VolumeBarSampler::new(*threshold as u64); for tick in ticks { if let Some(bar) = sampler.update(tick.price, tick.volume, tick.timestamp) { bars.push(bar); } } }, BarSamplingMethod::DollarBars(threshold) => { let mut sampler = DollarBarSampler::new(*threshold); for tick in ticks { if let Some(bar) = sampler.update(tick.price, tick.volume, tick.timestamp) { bars.push(bar); } } }, BarSamplingMethod::ImbalanceBars(threshold) => { if let Some(first_tick) = ticks.first() { // Initialize with first tick's price and timestamp let mut sampler = ImbalanceBarSampler::new( first_tick.price, *threshold, first_tick.timestamp, ); for tick in ticks { if let Some(bar) = sampler.update(tick.price, tick.volume, tick.timestamp) { bars.push(bar); } } } }, BarSamplingMethod::RunBars(threshold) => { let mut sampler = RunBarSampler::new(*threshold); for tick in ticks { if let Some(bar) = sampler.update(tick.price, tick.volume, tick.timestamp) { bars.push(bar); } } }, } Ok(bars) } /// Get symbol from processed message fn get_message_symbol(&self, msg: &ProcessedMessage) -> String { match msg { ProcessedMessage::Trade { symbol, .. } => symbol.clone(), ProcessedMessage::Quote { symbol, .. } => symbol.clone(), ProcessedMessage::OrderBook { symbol, .. } => symbol.clone(), ProcessedMessage::Ohlcv { symbol, .. } => symbol.clone(), ProcessedMessage::Status { .. } => "UNKNOWN".to_string(), } } /// Compute feature statistics for normalization fn compute_stats( &mut self, symbol_messages: &HashMap>, ) -> Result<()> { let mut prices = Vec::new(); let mut volumes = Vec::new(); for messages in symbol_messages.values() { 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)); }, _ => {}, } } } if prices.is_empty() { return Err(anyhow::anyhow!("No price data found for normalization")); } // Compute mean and std let price_mean = prices.iter().sum::() / prices.len() as f64; let price_var = prices.iter().map(|p| (p - price_mean).powi(2)).sum::() / prices.len() as f64; let price_std = price_var.sqrt().max(1e-8); let volume_mean = volumes.iter().sum::() / volumes.len() as f64; let volume_var = volumes .iter() .map(|v| (v - volume_mean).powi(2)) .sum::() / 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(()) } /// Create sequences from message list /// /// Uses sliding window with stride and optional max limit to prevent memory overflow. /// For 665K bars: stride=10, max=10K → 10K sequences instead of 665K fn create_sequences(&mut self, messages: &[ProcessedMessage]) -> Result> { let mut sequences = Vec::new(); // REFACTOR (Wave 5 Agent 26): Use production extract_ml_features() pipeline // This eliminates 43 zero-padded features (19.1% junk data) // Step 1: Convert ProcessedMessage to OHLCVBar format for production pipeline let bars: Vec = messages .iter() .filter_map(|msg| self.convert_message_to_bar(msg)) .collect(); if bars.is_empty() { return Err(anyhow::anyhow!("No OHLCV bars found in messages")); } // Step 2: Extract all features using production pipeline (requires 50-bar warmup) const WARMUP_PERIOD: usize = 50; if bars.len() < WARMUP_PERIOD + self.seq_len + 1 { return Err(anyhow::anyhow!( "Insufficient bars: {} available, need {} (warmup) + {} (seq_len) + 1 (target)", bars.len(), WARMUP_PERIOD, self.seq_len )); } debug!( "Extracting features using production pipeline: {} bars → 54 features per bar", bars.len() ); // extract_ml_features() returns Vec<[f64; 54]> after warmup period let feature_vectors = extract_ml_features(&bars) .context("Failed to extract 54-feature vectors from production pipeline")?; debug!( "✓ Extracted {} feature vectors (54 dims each) after warmup", feature_vectors.len() ); // Step 3: Create sequences from feature vectors // Calculate maximum possible sequences let max_possible = feature_vectors.len().saturating_sub(self.seq_len); // Apply stride (e.g., every 10th bar) let num_sequences_with_stride = (max_possible + self.stride - 1) / self.stride; // Apply max limit if specified let target_num_sequences = match self.max_sequences_per_symbol { Some(max) => num_sequences_with_stride.min(max), None => num_sequences_with_stride, }; debug!( "Sequence generation: {} feature vectors → {} sequences (stride={}, max={:?})", feature_vectors.len(), target_num_sequences, self.stride, self.max_sequences_per_symbol ); // Pre-allocate to prevent reallocation during loop sequences.reserve(target_num_sequences); // Progress tracking for large datasets let progress_interval = (target_num_sequences / 10).max(1); // Log every 10% // Sliding window over feature vectors with stride let mut seq_count = 0; let mut i = 0; while i < max_possible && seq_count < target_num_sequences { // Log progress every 10% if seq_count > 0 && seq_count % progress_interval == 0 { let progress = (seq_count as f64 / target_num_sequences as f64 * 100.0) as usize; debug!( " Sequence generation: {}% ({}/{})", progress, seq_count, target_num_sequences ); } // Extract seq_len feature vectors for input let mut features = Vec::with_capacity(self.seq_len * 54); for j in 0..self.seq_len { let feature_vec = &feature_vectors[i + j]; // Convert [f64; 54] to f32 and apply normalization let mut features_f32: Vec = feature_vec.iter().map(|&x| x as f32).collect(); features_f32 = self.normalize_features(&features_f32)?; features.extend_from_slice(&features_f32); } // FIXED (Agent 254): Target is next close price (regression), not full feature vector // Extract target price from the bar that corresponds to feature_vectors[i + seq_len] // Note: bars[warmup + i + seq_len] is the actual bar for feature_vectors[i + seq_len] let target_bar_idx = WARMUP_PERIOD + i + self.seq_len; let target_price = bars[target_bar_idx].close; // Normalize target price using same stats as features let normalized_target = ((target_price - self.stats.price_mean) / self.stats.price_std) as f32; // Create tensors with batch dimension // Input: [batch=1, seq_len, d_model] = [1, 60, 54] // Target: [batch=1, 1, 1] = single price for regression let input = Tensor::from_slice(&features, (1, self.seq_len, 54), &self.device)? .to_dtype(DType::F64)?; let target_tensor = Tensor::from_slice(&[normalized_target], (1, 1, 1), &self.device)? .to_dtype(DType::F64)?; sequences.push((input, target_tensor)); seq_count += 1; i += self.stride; // Apply stride (skip bars) } debug!( "✓ Generated {} sequences from {} feature vectors (0% zero-padding)", sequences.len(), feature_vectors.len() ); Ok(sequences) } /// Convert ProcessedMessage to OHLCVBar format for production feature extraction /// /// REFACTOR (Wave 5 Agent 26): Helper method to bridge DbnSequenceLoader with production pipeline fn convert_message_to_bar(&self, msg: &ProcessedMessage) -> Option { match msg { ProcessedMessage::Ohlcv { timestamp, open, high, low, close, volume, .. } => { // Convert HardwareTimestamp (nanos) to chrono::DateTime let nanos = timestamp.nanos; let secs = (nanos / 1_000_000_000) as i64; let nsec = (nanos % 1_000_000_000) as u32; let timestamp_dt = chrono::DateTime::from_timestamp(secs, nsec).unwrap_or_else(chrono::Utc::now); Some(OHLCVBar { timestamp: timestamp_dt, open: open.to_f64(), high: high.to_f64(), low: low.to_f64(), close: close.to_f64(), volume: volume.to_f64().unwrap_or(0.0), }) }, // Only OHLCV messages are supported for production feature extraction // Trade and Quote messages are not bar-based _ => None, } } /// Extract target price (close price) for regression /// /// FIXED (Agent 254): Model output_dim=1 for price prediction (regression) /// Target should be single close price, not full 256-dim feature vector fn extract_target_price(&self, msg: &ProcessedMessage) -> Result { match msg { ProcessedMessage::Ohlcv { close, .. } => { // Normalize close price using same stats as features let c = (close.to_f64() - self.stats.price_mean) / self.stats.price_std; Ok(c as f32) }, ProcessedMessage::Trade { price, .. } => { // For trades, use price as target let p = (price.to_f64() - self.stats.price_mean) / self.stats.price_std; Ok(p as f32) }, ProcessedMessage::Quote { ask, bid, .. } => { // For quotes, use mid-price as target let mid = match (ask, bid) { (Some(a), Some(b)) => (a.to_f64() + b.to_f64()) / 2.0, (Some(a), None) => a.to_f64(), (None, Some(b)) => b.to_f64(), _ => 0.0, }; let normalized = (mid - self.stats.price_mean) / self.stats.price_std; Ok(normalized as f32) }, _ => { // Default: return 0.0 for other message types Ok(0.0) }, } } /// Extract normalized features from a message (dynamic based on FeatureConfig) /// /// FIXED (Agent C2): Removed 54-feature padding bug. Now extracts real features /// based on FeatureConfig (Wave A: 26, Wave B: 36, Wave C: 65+). /// /// Wave A features (26): /// LEGACY METHOD (Wave 5 Agent 26): This method is DEPRECATED and should NOT be used. /// /// Use production `extract_ml_features()` pipeline instead (0% zero-padding). /// /// This method produced 43 zero-padded features (19.1% junk data): /// - Alternative bar features (10): lines 1221-1227 /// - Microstructure features (3): lines 1230-1236 /// - Fractional differentiation (20): lines 1239-1244 /// - Regime detection (10): lines 1247-1252 /// /// ## Original Feature Breakdown (Wave A only - 26 features): /// - Base OHLCV (5 features): open, high, low, close, volume /// - Derived features (4 features): range, body, upper_wick, lower_wick /// - Price ratios (10 features): c/o, h/l, h/c, l/c, c/h, c/l, body/range, upper_wick/range, lower_wick/range, v/price /// - Log returns (4 features): ln(c/o), ln(h/o), ln(l/o), ln(c/h) /// - Price deltas (3 features): c-o, h-o, l-o (removed c-l to match 26 total) /// /// Total: 5 + 4 + 10 + 4 + 3 = 26 features (matches FeatureConfig::wave_a()) #[allow(dead_code)] fn extract_features(&mut self, msg: &ProcessedMessage) -> Result> { match msg { ProcessedMessage::Ohlcv { open, high, low, close, volume, .. } => { // Normalize OHLCV 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; // High-Low range let body = c - o; // Close-Open (candle body) let upper_wick = h - c.max(o); // Upper shadow let lower_wick = l.min(o) - l; // Lower shadow // Build feature vector based on FeatureConfig let mut features = Vec::with_capacity(self.d_model); // 1. Base OHLCV (5 features) - always included in Wave A/B/C if self.feature_config.enable_ohlcv { features.push(o as f32); features.push(h as f32); features.push(l as f32); features.push(c as f32); features.push(v as f32); } // 2. Derived features (4 features) - part of Wave A technical indicators if self.feature_config.enable_technical_indicators { features.push(range as f32); features.push(body as f32); features.push(upper_wick as f32); features.push(lower_wick as f32); } // 3. Price ratios (10 features) - part of Wave A technical indicators if self.feature_config.enable_technical_indicators { let safe_div = |a: f64, b: f64| if b.abs() > 1e-8 { (a / b) as f32 } else { 0.0 }; features.push(safe_div(c, o)); // close/open ratio features.push(safe_div(h, l)); // high/low ratio features.push(safe_div(h, c)); // high/close ratio features.push(safe_div(l, c)); // low/close ratio features.push(safe_div(c, h)); // close/high ratio (upper position) features.push(safe_div(c, l)); // close/low ratio (lower position) features.push(safe_div(body.abs(), range.max(1e-8))); // body/range ratio features.push(safe_div(upper_wick, range.max(1e-8))); // upper wick ratio features.push(safe_div(lower_wick, range.max(1e-8))); // lower wick ratio features.push(safe_div(v, (h + l + c + o) / 4.0)); // volume/price ratio } // 4. Log returns (4 features) - part of Wave A technical indicators if self.feature_config.enable_technical_indicators { let safe_ln = |a: f64, b: f64| { let ratio = a / b.max(1e-8); if ratio > 0.0 { ratio.ln() as f32 } else { 0.0 // Return 0 for negative or zero ratios (normalized prices can be negative) } }; features.push(safe_ln(c, o)); // log return features.push(safe_ln(h, o)); // log high return features.push(safe_ln(l, o)); // log low return features.push(safe_ln(c, h)); // log close/high } // 5. Price deltas (3 features) - part of Wave A technical indicators // REMOVED: (c - l) to match 26-feature count if self.feature_config.enable_technical_indicators { features.push((c - o) as f32); // raw price change features.push((h - o) as f32); // open to high features.push((l - o) as f32); // open to low } // 6. Alternative bar features (10 features) - Wave B if self.feature_config.enable_alternative_bars { // TODO (Wave B): Add dollar bar, volume bar, tick bar, run bar, imbalance bar features // For now, pad with zeros for _ in 0..10 { features.push(0.0); } } // 7. Microstructure features (3 features) - Wave A/C if self.feature_config.enable_microstructure { // TODO: Add Amihud Illiquidity, Roll Measure, Corwin-Schultz Spread // For now, pad with zeros (not yet integrated) for _ in 0..3 { features.push(0.0); } } // 8. Fractional differentiation features (20 features) - Wave C if self.feature_config.enable_fractional_diff { // TODO (Wave C): Add fractional differentiation features for _ in 0..20 { features.push(0.0); } } // 9. Regime detection features (10 features) - Wave C if self.feature_config.enable_regime_detection { // TODO (Wave C): Add CUSUM structural breaks, regime indicators for _ in 0..10 { features.push(0.0); } } // 10. Wave D regime features (24 features) - Wave D if self.feature_config.enable_wave_d_regime { // CUSUM Statistics (indices 201-210, 10 features) // Use log return as input to CUSUM detector let log_return = if o.abs() > 1e-8 { (c / o).ln() } else { 0.0 }; let cusum_features = self.regime_cusum.update(log_return); for &feat in &cusum_features { features.push(feat as f32); } // ADX & Directional Indicators (indices 211-215, 5 features) // Create OHLCVBar for ADX calculation let current_bar_adx = OHLCVBar { timestamp: chrono::DateTime::::from_timestamp(0, 0).unwrap(), // Not used in feature calculation open: open.to_f64(), high: high.to_f64(), low: low.to_f64(), close: close.to_f64(), volume: volume.to_f64().unwrap_or(0.0), }; // Add current bar to ADX buffer self.bar_buffer_adx.push(current_bar_adx.clone()); if self.bar_buffer_adx.len() > 100 { self.bar_buffer_adx.remove(0); // Keep last 100 bars } let adx_features = self.regime_adx.update(¤t_bar_adx); for &feat in &adx_features { features.push(feat as f32); } // Regime Transition Probabilities (indices 216-220, 5 features) // Update transition matrix with current regime and compute features let transition_features = self.regime_transition.update(self.current_regime); for &feat in &transition_features { features.push(feat as f32); } // Adaptive Strategy Metrics (indices 221-224, 4 features) // Create OHLCVBar for adaptive features (requires DateTime timestamp) let current_bar_adaptive = OHLCVBar { timestamp: chrono::Utc::now(), // Use current time open: open.to_f64(), high: high.to_f64(), low: low.to_f64(), close: close.to_f64(), volume: volume.to_f64().unwrap_or(0.0), }; // Add current bar to adaptive buffer self.bar_buffer_adaptive.push(current_bar_adaptive); if self.bar_buffer_adaptive.len() > 100 { self.bar_buffer_adaptive.remove(0); // Keep last 100 bars } // Use recent returns and current position size let adaptive_features = self.regime_adaptive.update( self.current_regime, log_return, 50_000.0, // Current position size (50% of max) &self.bar_buffer_adaptive, ); for &feat in &adaptive_features { features.push(feat as f32); } } // Sanity check: ensure feature count matches FeatureConfig debug_assert_eq!( features.len(), self.feature_config.feature_count(), "Feature vector must match FeatureConfig.feature_count(): expected {}, got {}", self.feature_config.feature_count(), features.len() ); Ok(features) }, ProcessedMessage::Trade { price, size, .. } => { // Trade messages: create feature vector based on FeatureConfig let p = (price.to_f64() - self.stats.price_mean) / self.stats.price_std; let s = (size.to_f64().unwrap_or(0.0) - self.stats.volume_mean) / self.stats.volume_std; // Use price as OHLCV (all same for trades) let mut features = Vec::with_capacity(self.d_model); if self.feature_config.enable_ohlcv { features.push(p as f32); // open features.push(p as f32); // high features.push(p as f32); // low features.push(p as f32); // close features.push(s as f32); // volume } // Pad remaining features with zeros while features.len() < self.feature_config.feature_count() { features.push(0.0); } Ok(features) }, ProcessedMessage::Quote { bid, ask, .. } => { // Quote messages: create feature vector based on FeatureConfig let b = bid .map(|p| (p.to_f64() - self.stats.price_mean) / self.stats.price_std) .unwrap_or(0.0); let a = ask .map(|p| (p.to_f64() - self.stats.price_mean) / self.stats.price_std) .unwrap_or(0.0); let spread = a - b; let mid = (a + b) / 2.0; // Use mid-price as OHLCV (all same for quotes) let mut features = Vec::with_capacity(self.d_model); if self.feature_config.enable_ohlcv { features.push(mid as f32); // open features.push(a as f32); // high (ask) features.push(b as f32); // low (bid) features.push(mid as f32); // close features.push(spread as f32); // volume (use spread) } // Pad remaining features with zeros while features.len() < self.feature_config.feature_count() { features.push(0.0); } Ok(features) }, _ => { // Default fallback: zero vector of d_model dimensions Ok(vec![0.0; self.feature_config.feature_count()]) }, } } /// Normalize features using FeatureNormalizer (Agent F1: Critical for numerical stability) /// /// Converts f32 features to f64, applies normalization, and converts back to f32. /// This prevents numerical instability in MAMBA-2 training (loss values at 10^38 scale). /// /// # Arguments /// * `features` - Raw features (26/36/65/54 dimensions) /// /// # Returns /// Normalized features with all values in reasonable ranges fn normalize_features(&self, features: &[f32]) -> Result> { // Convert f32 -> f64 (FeatureNormalizer uses f64) let mut feature_vec_f64: [f64; 51] = [0.0; 51]; for (i, &val) in features.iter().enumerate() { if i < 51 { feature_vec_f64[i] = val as f64; } } // Apply normalization (requires mutable borrow) // We'll use interior mutability pattern with Cell/RefCell in production // For now, we'll compute normalized features without state update // (warmup phase will be skipped, but normalization logic still applies) // Instead, let's compute normalization manually based on feature indices // This avoids the mutable borrow issue while still providing normalization self.apply_manual_normalization(&mut feature_vec_f64)?; // Convert back to f32 let mut normalized_f32 = Vec::with_capacity(features.len()); for i in 0..features.len() { normalized_f32.push(feature_vec_f64[i] as f32); } Ok(normalized_f32) } /// Apply manual normalization based on feature indices (Agent F1) /// /// This is a stateless normalization that applies scaling without requiring /// rolling window state updates. Uses fixed scaling factors appropriate for /// each feature category. fn apply_manual_normalization(&self, features: &mut [f64; 51]) -> Result<()> { // Skip OHLCV (indices 0-4): already normalized by extract_features() // Skip technical indicators (indices 5-14): already in normalized ranges // Normalize price features (indices 15-74): z-score with clipping for i in 15..75 { if i < features.len() { // Clip to ±3σ range features[i] = features[i].clamp(-3.0, 3.0); } } // Normalize volume features (indices 75-114): percentile rank [0, 1] for i in 75..115 { if i < features.len() { // Clip to [0, 1] range features[i] = features[i].clamp(0.0, 1.0); } } // Normalize microstructure features (indices 115-164): log+z-score for i in 115..165 { if i < features.len() { // Clip to ±3σ range features[i] = features[i].clamp(-3.0, 3.0); } } // Skip time/statistical features (indices 165-200): already normalized // Normalize Wave D CUSUM features (indices 201-210): z-score for i in 201..211 { if i < features.len() { features[i] = features[i].clamp(-3.0, 3.0); } } // Normalize Wave D ADX features (indices 211-215): [0, 1] for i in 211..216 { if i < features.len() { features[i] = features[i].clamp(0.0, 1.0); } } // Normalize Wave D transition features (indices 216-220): z-score for i in 216..221 { if i < features.len() { features[i] = features[i].clamp(-3.0, 3.0); } } // Normalize Wave D adaptive features (indices 221-224): [0, 2] for i in 50..54 { if i < features.len() { features[i] = features[i].clamp(0.0, 2.0); } } // Final validation: ensure all features are finite for (i, &val) in features.iter().enumerate() { if !val.is_finite() { anyhow::bail!("Feature {} is non-finite after normalization: {}", i, val); } } Ok(()) } } #[cfg(test)] mod tests { use super::*; #[tokio::test] async fn test_loader_creation_wave_a() { // Wave A: 26 features let loader = DbnSequenceLoader::new(60, 26).await; assert!(loader.is_ok()); let loader = loader.unwrap(); assert_eq!(loader.d_model, 26); assert_eq!(loader.feature_config.feature_count(), 26); } #[tokio::test] async fn test_loader_with_feature_config_wave_b() { // Wave B: 36 features let config = crate::features::config::FeatureConfig::wave_b(); let loader = DbnSequenceLoader::with_feature_config(60, config).await; assert!(loader.is_ok()); let loader = loader.unwrap(); assert_eq!(loader.d_model, 36); assert_eq!(loader.feature_config.feature_count(), 36); } #[tokio::test] async fn test_loader_with_feature_config_wave_c() { // Wave C: 65+ features let config = crate::features::config::FeatureConfig::wave_c(); let loader = DbnSequenceLoader::with_feature_config(60, config).await; assert!(loader.is_ok()); let loader = loader.unwrap(); assert!(loader.d_model >= 65); assert_eq!(loader.feature_config.feature_count(), loader.d_model); } #[tokio::test] async fn test_loader_rejects_mismatched_d_model() { // Should fail: d_model=256 does not match Wave A (26 features) let loader = DbnSequenceLoader::new(60, 256).await; assert!(loader.is_err()); let err = loader.unwrap_err(); assert!(err.to_string().contains("does not match")); } #[test] fn test_feature_stats_default() { let stats = FeatureStats::default(); assert_eq!(stats.price_mean, 0.0); assert_eq!(stats.price_std, 1.0); } }