ZN.FUT tests in crates/ml/src/data_loader.rs were failing because no
valid ZN.FUT DBN data is available locally; gated with #[ignore].
ml-asset-selection's universe definition and backtesting's
zn_futures() slippage profile remain untouched — those are production
references to ZN as a candidate symbol, distinct from data availability.
Migrated 4 deprecated cudarc memcpy_stod/memcpy_dtov sites in the
test function test_eval_action_select_eval_argmax_picks_best in
crates/ml/src/cuda_pipeline/mod.rs to mapped-pinned per
feedback_no_htod_htoh_only_mapped_pinned:
- 3x memcpy_stod (f32 input uploads) → MappedF32Buffer::new +
write_from_slice + dev_ptr as raw u64 kernel arg; kernel reads
directly from mapped-pinned pages, no DtoD copy needed
- 1x memcpy_dtov (i32 output readback) → MappedI32Buffer::new +
dev_ptr as kernel arg + read_all() after stream sync
The cudarc deprecation suggested clone_htod/clone_dtoh as replacements
but those still perform HtoD/DtoH copies — violating the strict rule.
Mapped-pinned with direct dev_ptr kernel args is the correct pattern
(matches distributional_q_tests.rs).
Note: DqnGpuData/PpoGpuData upload paths also in mod.rs still use
clone_to_device_f32_via_pinned; migrating those requires changing
CudaSlice<f32> struct fields to MappedF32Buffer which is blocked until
gpu_dqn_trainer.rs consumers are also updated (separate scope).
Workspace cargo check warnings: 15 → 15 (test-only deprecated calls
not visible to cargo check; ZN gate adds 3 to ignored count).
cargo test -p ml --lib failures: 16 → 13 (3 ZN tests now ignored).
Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
745 lines
26 KiB
Rust
745 lines
26 KiB
Rust
//! # Real Data Loader - DBN to ML Features
|
|
//!
|
|
//! Loads real market data from Databento DBN files and converts to ML-ready features.
|
|
//! This module bridges the gap between raw market data and ML model input, providing
|
|
//! feature extraction, technical indicators, and data quality validation.
|
|
//!
|
|
//! ## Architecture
|
|
//!
|
|
//! ```text
|
|
//! ┌──────────────┐ ┌──────────────┐ ┌──────────────┐
|
|
//! │ DBN Files │────▶│ Data Loader │────▶│ ML Features │
|
|
//! │ (Parquet) │ │ + Indicators│ │ (Normalized)│
|
|
//! └──────────────┘ └──────────────┘ └──────────────┘
|
|
//! │ │ │
|
|
//! ▼ ▼ ▼
|
|
//! OHLCV Bars Technical Indicators Feature Matrix
|
|
//! Timestamps RSI, MACD, BB, etc. Ready for ML
|
|
//! ```
|
|
//!
|
|
//! ## Usage
|
|
//!
|
|
//! ```rust
|
|
//! use ml::data_loader::RealDataLoader;
|
|
//!
|
|
//! let loader = RealDataLoader::new("test_data/real/databento").await?;
|
|
//! let bars = loader.load_symbol_data("ZN.FUT").await?;
|
|
//! let features = loader.extract_features(&bars)?;
|
|
//! let indicators = loader.calculate_indicators(&bars)?;
|
|
//! ```
|
|
|
|
use anyhow::{Context, Result};
|
|
use chrono::DateTime;
|
|
use dbn::decode::{DbnDecoder, DecodeRecordRef};
|
|
use dbn::OhlcvMsg;
|
|
use std::collections::HashMap;
|
|
use std::path::{Path, PathBuf};
|
|
use tracing::{debug, info};
|
|
|
|
use crate::types::OHLCVBar;
|
|
|
|
/// Feature matrix for ML model input
|
|
///
|
|
/// Contains normalized features ready for ML training/inference:
|
|
/// - prices: OHLCV data (normalized)
|
|
/// - returns: Log returns
|
|
/// - volume: Normalized volume
|
|
/// - indicators: Technical indicators (10 essential ones)
|
|
#[derive(Debug, Clone)]
|
|
pub struct FeatureMatrix {
|
|
/// OHLCV prices (each bar is [open, high, low, close, volume])
|
|
pub prices: Vec<Vec<f32>>,
|
|
/// Log returns (close-to-close)
|
|
pub returns: Vec<f32>,
|
|
/// Normalized volume
|
|
pub volume: Vec<f32>,
|
|
/// Technical indicators
|
|
pub indicators: Vec<Vec<f32>>,
|
|
}
|
|
|
|
/// Technical indicators (10 essential ones)
|
|
///
|
|
/// All indicators are calculated with standard parameters for 1-minute OHLCV data:
|
|
/// - RSI(14): Relative Strength Index
|
|
/// - MACD(12,26,9): Moving Average Convergence Divergence
|
|
/// - Bollinger Bands(20, 2.0): Price envelope
|
|
/// - ATR(14): Average True Range
|
|
/// - EMA(12, 26): Exponential moving averages
|
|
/// - Volume MA(20): Volume moving average
|
|
#[derive(Debug, Clone)]
|
|
pub struct Indicators {
|
|
/// RSI(14) - values 0-100
|
|
pub rsi: Vec<f32>,
|
|
/// MACD line (12,26)
|
|
pub macd: Vec<f32>,
|
|
/// MACD signal line (9)
|
|
pub macd_signal: Vec<f32>,
|
|
/// Bollinger upper band (20, 2.0)
|
|
pub bb_upper: Vec<f32>,
|
|
/// Bollinger middle band (SMA 20)
|
|
pub bb_middle: Vec<f32>,
|
|
/// Bollinger lower band (20, 2.0)
|
|
pub bb_lower: Vec<f32>,
|
|
/// ATR(14) - volatility measure
|
|
pub atr: Vec<f32>,
|
|
/// EMA(12) - fast exponential moving average
|
|
pub ema_fast: Vec<f32>,
|
|
/// EMA(26) - slow exponential moving average
|
|
pub ema_slow: Vec<f32>,
|
|
/// Volume MA(20) - volume moving average
|
|
pub volume_ma: Vec<f32>,
|
|
}
|
|
|
|
/// Real data loader for DBN files
|
|
///
|
|
/// Loads OHLCV data from Databento DBN files and extracts ML-ready features.
|
|
#[derive(Debug)]
|
|
pub struct RealDataLoader {
|
|
/// Base directory containing DBN files
|
|
base_path: PathBuf,
|
|
/// Cached symbol data
|
|
cache: HashMap<String, Vec<OHLCVBar>>,
|
|
}
|
|
|
|
impl RealDataLoader {
|
|
/// Create new data loader with automatic workspace root detection
|
|
///
|
|
/// Checks `FOXHUNT_DATA_DIR` env var first (used in CI with PVC mount),
|
|
/// then falls back to finding workspace root by looking for Cargo.toml.
|
|
pub fn new_from_workspace() -> Result<Self> {
|
|
// CI / PVC override: FOXHUNT_DATA_DIR points directly at the data directory
|
|
if let Ok(data_dir) = std::env::var("FOXHUNT_DATA_DIR") {
|
|
let base_path = PathBuf::from(&data_dir);
|
|
if base_path.exists() {
|
|
info!("Using FOXHUNT_DATA_DIR: {}", data_dir);
|
|
return Ok(Self {
|
|
base_path,
|
|
cache: HashMap::new(),
|
|
});
|
|
}
|
|
}
|
|
|
|
// CI: TEST_DATA_DIR points to test-data PVC (e.g. /data/test-data).
|
|
// DBN files live under ohlcv/ subdirectory with per-symbol layout.
|
|
if let Ok(data_dir) = std::env::var("TEST_DATA_DIR") {
|
|
let ohlcv = PathBuf::from(&data_dir).join("ohlcv");
|
|
if ohlcv.is_dir() {
|
|
info!("Using TEST_DATA_DIR/ohlcv: {}", ohlcv.display());
|
|
return Ok(Self {
|
|
base_path: ohlcv,
|
|
cache: HashMap::new(),
|
|
});
|
|
}
|
|
// Flat layout fallback
|
|
let base_path = PathBuf::from(&data_dir);
|
|
if base_path.exists() {
|
|
info!("Using TEST_DATA_DIR: {}", data_dir);
|
|
return Ok(Self {
|
|
base_path,
|
|
cache: HashMap::new(),
|
|
});
|
|
}
|
|
}
|
|
|
|
let mut current = std::env::current_dir()?;
|
|
|
|
// Try to find workspace root
|
|
while !current.join("Cargo.toml").exists() || !current.join("test_data").exists() {
|
|
if !current.pop() {
|
|
return Err(anyhow::anyhow!("Could not find workspace root"));
|
|
}
|
|
}
|
|
|
|
// Try paths in order: futures-baseline (primary), real/databento (legacy)
|
|
let candidates = [
|
|
current.join("test_data/futures-baseline"),
|
|
current.join("data/cache/futures-baseline"),
|
|
current.join("test_data/real/databento"),
|
|
];
|
|
let base_path = candidates.iter()
|
|
.find(|p| p.is_dir())
|
|
.cloned()
|
|
.unwrap_or_else(|| current.join("test_data/futures-baseline"));
|
|
Ok(Self {
|
|
base_path,
|
|
cache: HashMap::new(),
|
|
})
|
|
}
|
|
|
|
/// Create new data loader
|
|
///
|
|
/// # Arguments
|
|
///
|
|
/// * `base_path` - Directory containing DBN files (e.g., "test_data/real/databento")
|
|
pub fn new<P: AsRef<Path>>(base_path: P) -> Self {
|
|
Self {
|
|
base_path: base_path.as_ref().to_path_buf(),
|
|
cache: HashMap::new(),
|
|
}
|
|
}
|
|
|
|
/// Load OHLCV data from DBN file
|
|
///
|
|
/// Searches for DBN files matching the symbol pattern and loads OHLCV bars.
|
|
/// Caches loaded data for repeated access.
|
|
///
|
|
/// # Arguments
|
|
///
|
|
/// * `symbol` - Symbol to load (e.g., "ZN.FUT", "6E.FUT", "ES.FUT")
|
|
///
|
|
/// # Returns
|
|
///
|
|
/// Vector of OHLCV bars sorted by timestamp
|
|
pub async fn load_symbol_data(&mut self, symbol: &str) -> Result<Vec<OHLCVBar>> {
|
|
// Check cache first
|
|
if let Some(cached) = self.cache.get(symbol) {
|
|
debug!(
|
|
"Returning cached data for {}: {} bars",
|
|
symbol,
|
|
cached.len()
|
|
);
|
|
return Ok(cached.clone());
|
|
}
|
|
|
|
info!("Loading DBN data for symbol: {}", symbol);
|
|
|
|
// Find DBN file for this symbol
|
|
let dbn_file = self.find_dbn_file(symbol)?;
|
|
info!("Found DBN file: {:?}", dbn_file);
|
|
|
|
// Load and parse DBN file
|
|
let bars = self.parse_dbn_file(&dbn_file)?;
|
|
info!("Loaded {} bars for {}", bars.len(), symbol);
|
|
|
|
// Cache the data
|
|
self.cache.insert(symbol.to_string(), bars.clone());
|
|
|
|
Ok(bars)
|
|
}
|
|
|
|
/// Find DBN file for symbol
|
|
///
|
|
/// Searches base_path and per-symbol subdirectories for files matching
|
|
/// `{symbol}*.dbn` or `{symbol}*.dbn.zst`. Preference order:
|
|
/// 1. `*.uncompressed.dbn` (pre-decompressed)
|
|
/// 2. `*.dbn` (raw uncompressed)
|
|
/// 3. `*.dbn.zst` (zstd-compressed, standard Databento download format)
|
|
fn find_dbn_file(&self, symbol: &str) -> Result<PathBuf> {
|
|
// Search base_path and also the per-symbol subdirectory (Databento layout)
|
|
let mut search_dirs = vec![self.base_path.clone()];
|
|
let subdir = self.base_path.join(symbol);
|
|
if subdir.is_dir() {
|
|
search_dirs.push(subdir);
|
|
}
|
|
|
|
let mut uncompressed = Vec::new();
|
|
let mut compressed = Vec::new();
|
|
|
|
for search_dir in &search_dirs {
|
|
let dir = match std::fs::read_dir(search_dir) {
|
|
Ok(d) => d,
|
|
Err(_) => continue,
|
|
};
|
|
|
|
for entry in dir {
|
|
let entry = entry?;
|
|
let path = entry.path();
|
|
let filename = path.file_name().and_then(|s| s.to_str()).unwrap_or("");
|
|
|
|
if !filename.starts_with(symbol) {
|
|
continue;
|
|
}
|
|
|
|
// Prefer uncompressed files (dbn 0.42.0 compatibility)
|
|
if filename.contains(".uncompressed.dbn") {
|
|
return Ok(path);
|
|
}
|
|
if filename.ends_with(".dbn.zst") {
|
|
compressed.push(path);
|
|
} else if filename.ends_with(".dbn") {
|
|
uncompressed.push(path);
|
|
} else {
|
|
// Skip non-DBN files
|
|
}
|
|
}
|
|
}
|
|
|
|
// Prefer raw .dbn over .dbn.zst
|
|
if let Some(path) = uncompressed.into_iter().next() {
|
|
return Ok(path);
|
|
}
|
|
compressed
|
|
.into_iter()
|
|
.next()
|
|
.ok_or_else(|| anyhow::anyhow!("No DBN file found for symbol: {}", symbol))
|
|
}
|
|
|
|
/// Parse DBN file and extract OHLCV bars
|
|
///
|
|
/// Uses the `dbn` crate to decode DBN binary format and extract OHLCV records.
|
|
/// Handles both raw `.dbn` and zstd-compressed `.dbn.zst` files.
|
|
fn parse_dbn_file(&self, path: &Path) -> Result<Vec<OHLCVBar>> {
|
|
// Detect zstd by magic bytes — file extension is unreliable
|
|
// (Databento downloads are often zstd-compressed with .dbn extension)
|
|
let bars = if Self::is_zstd_file(path)? {
|
|
let mut decoder = DbnDecoder::from_zstd_file(path)
|
|
.context(format!("Failed to open zstd DBN file: {:?}", path))?;
|
|
Self::decode_ohlcv_bars(&mut decoder)?
|
|
} else {
|
|
let mut decoder = DbnDecoder::from_file(path)
|
|
.context(format!("Failed to open DBN file: {:?}", path))?;
|
|
Self::decode_ohlcv_bars(&mut decoder)?
|
|
};
|
|
|
|
Ok(bars)
|
|
}
|
|
|
|
/// Detect zstd-compressed files by magic bytes (0x28B52FFD).
|
|
fn is_zstd_file(path: &Path) -> Result<bool> {
|
|
let mut file = std::fs::File::open(path)?;
|
|
let mut magic = [0_u8; 4];
|
|
use std::io::Read;
|
|
if file.read_exact(&mut magic).is_ok() {
|
|
Ok(magic == [0x28, 0xB5, 0x2F, 0xFD])
|
|
} else {
|
|
Ok(false)
|
|
}
|
|
}
|
|
|
|
/// Decode OHLCV bars from a DBN decoder (generic over reader type)
|
|
fn decode_ohlcv_bars<R: std::io::Read>(decoder: &mut DbnDecoder<R>) -> Result<Vec<OHLCVBar>> {
|
|
let mut bars = Vec::new();
|
|
|
|
// Iterate over records
|
|
while let Some(record_ref) = decoder
|
|
.decode_record_ref()
|
|
.context("Failed to decode record")?
|
|
{
|
|
if let Some(ohlcv) = record_ref.get::<OhlcvMsg>() {
|
|
// Convert DBN record to OHLCVBar
|
|
let bar = OHLCVBar {
|
|
timestamp: DateTime::from_timestamp_nanos(ohlcv.hd.ts_event as i64),
|
|
open: ohlcv.open as f64 / 1e9, // DBN stores prices in fixed-point
|
|
high: ohlcv.high as f64 / 1e9,
|
|
low: ohlcv.low as f64 / 1e9,
|
|
close: ohlcv.close as f64 / 1e9,
|
|
volume: ohlcv.volume as f64,
|
|
};
|
|
|
|
bars.push(bar);
|
|
}
|
|
}
|
|
|
|
// Sort by timestamp (should already be sorted but ensure it)
|
|
bars.sort_by_key(|b| b.timestamp);
|
|
|
|
Ok(bars)
|
|
}
|
|
|
|
/// Extract basic features for ML models
|
|
///
|
|
/// Converts OHLCV bars to normalized feature matrix:
|
|
/// - prices: OHLCV data (normalized to 0-1 range per feature)
|
|
/// - returns: Log returns (close-to-close)
|
|
/// - volume: Normalized volume
|
|
///
|
|
/// # Arguments
|
|
///
|
|
/// * `bars` - OHLCV bars to extract features from
|
|
pub fn extract_features(&self, bars: &[OHLCVBar]) -> Result<FeatureMatrix> {
|
|
if bars.is_empty() {
|
|
return Err(anyhow::anyhow!(
|
|
"Cannot extract features from empty bar sequence"
|
|
));
|
|
}
|
|
|
|
let mut prices = Vec::with_capacity(bars.len());
|
|
let mut returns = Vec::with_capacity(bars.len());
|
|
let mut volume = Vec::with_capacity(bars.len());
|
|
|
|
// Calculate normalization factors
|
|
let (price_min, price_max) = self.price_range(bars);
|
|
let (vol_min, vol_max) = self.volume_range(bars);
|
|
|
|
// Extract features
|
|
for (i, bar) in bars.iter().enumerate() {
|
|
// Normalize OHLCV to 0-1 range
|
|
let norm_open = ((bar.open - price_min) / (price_max - price_min)) as f32;
|
|
let norm_high = ((bar.high - price_min) / (price_max - price_min)) as f32;
|
|
let norm_low = ((bar.low - price_min) / (price_max - price_min)) as f32;
|
|
let norm_close = ((bar.close - price_min) / (price_max - price_min)) as f32;
|
|
let norm_volume = ((bar.volume - vol_min) / (vol_max - vol_min)) as f32;
|
|
|
|
prices.push(vec![
|
|
norm_open,
|
|
norm_high,
|
|
norm_low,
|
|
norm_close,
|
|
norm_volume,
|
|
]);
|
|
volume.push(norm_volume);
|
|
|
|
// Calculate log returns (skip first bar)
|
|
if i > 0 {
|
|
let log_return = ((bar.close / bars[i - 1].close).ln()) as f32;
|
|
returns.push(log_return);
|
|
}
|
|
}
|
|
|
|
// First return is 0 (no previous bar)
|
|
if !returns.is_empty() {
|
|
returns.insert(0, 0.0);
|
|
}
|
|
|
|
Ok(FeatureMatrix {
|
|
prices,
|
|
returns,
|
|
volume,
|
|
indicators: Vec::new(), // Filled by calculate_indicators()
|
|
})
|
|
}
|
|
|
|
/// Calculate technical indicators
|
|
///
|
|
/// Computes 10 essential technical indicators:
|
|
/// - RSI(14): Relative Strength Index
|
|
/// - MACD(12,26,9): Moving Average Convergence Divergence
|
|
/// - Bollinger Bands(20, 2.0): Price envelope
|
|
/// - ATR(14): Average True Range
|
|
/// - EMA(12, 26): Exponential moving averages
|
|
/// - Volume MA(20): Volume moving average
|
|
///
|
|
/// # Arguments
|
|
///
|
|
/// * `bars` - OHLCV bars to calculate indicators from
|
|
pub fn calculate_indicators(&self, bars: &[OHLCVBar]) -> Result<Indicators> {
|
|
if bars.len() < 26 {
|
|
return Err(anyhow::anyhow!(
|
|
"Need at least 26 bars to calculate indicators (got {})",
|
|
bars.len()
|
|
));
|
|
}
|
|
|
|
let closes: Vec<f64> = bars.iter().map(|b| b.close).collect();
|
|
let volumes: Vec<f64> = bars.iter().map(|b| b.volume).collect();
|
|
|
|
Ok(Indicators {
|
|
rsi: self.calculate_rsi(&closes, 14)?,
|
|
macd: self.calculate_macd(&closes, 12, 26)?.iter().map(|&x| x as f32).collect(),
|
|
macd_signal: self.calculate_macd_signal(&closes, 12, 26, 9)?.iter().map(|&x| x as f32).collect(),
|
|
bb_upper: self.calculate_bb_upper(&closes, 20, 2.0)?,
|
|
bb_middle: self.calculate_sma(&closes, 20)?.iter().map(|&x| x as f32).collect(),
|
|
bb_lower: self.calculate_bb_lower(&closes, 20, 2.0)?,
|
|
atr: self.calculate_atr(bars, 14)?,
|
|
ema_fast: self.calculate_ema(&closes, 12)?.iter().map(|&x| x as f32).collect(),
|
|
ema_slow: self.calculate_ema(&closes, 26)?.iter().map(|&x| x as f32).collect(),
|
|
volume_ma: self.calculate_sma(&volumes, 20)?.iter().map(|&x| x as f32).collect(),
|
|
})
|
|
}
|
|
|
|
// ===== Technical Indicator Calculations =====
|
|
|
|
/// Calculate RSI (Relative Strength Index)
|
|
fn calculate_rsi(&self, prices: &[f64], period: usize) -> Result<Vec<f32>> {
|
|
let mut rsi = Vec::with_capacity(prices.len());
|
|
|
|
for i in 0..prices.len() {
|
|
if i < period {
|
|
rsi.push(50.0); // Neutral RSI for warmup period
|
|
continue;
|
|
}
|
|
|
|
let mut gains = 0.0;
|
|
let mut losses = 0.0;
|
|
|
|
for j in (i - period + 1)..=i {
|
|
let change = prices[j] - prices[j - 1];
|
|
if change > 0.0 {
|
|
gains += change;
|
|
} else {
|
|
losses += -change;
|
|
}
|
|
}
|
|
|
|
let avg_gain = gains / period as f64;
|
|
let avg_loss = losses / period as f64;
|
|
|
|
let rs = if avg_loss > 0.0 {
|
|
avg_gain / avg_loss
|
|
} else {
|
|
100.0 // Max RSI when no losses
|
|
};
|
|
|
|
let rsi_value = 100.0 - (100.0 / (1.0 + rs));
|
|
rsi.push(rsi_value as f32);
|
|
}
|
|
|
|
Ok(rsi)
|
|
}
|
|
|
|
/// Calculate MACD line
|
|
///
|
|
/// Returns f64 to avoid precision-losing f64→f32→f64 round-trips
|
|
/// when the result feeds back into `calculate_ema` for the signal line.
|
|
fn calculate_macd(&self, prices: &[f64], fast: usize, slow: usize) -> Result<Vec<f64>> {
|
|
let ema_fast = self.calculate_ema(prices, fast)?;
|
|
let ema_slow = self.calculate_ema(prices, slow)?;
|
|
|
|
Ok(ema_fast
|
|
.iter()
|
|
.zip(ema_slow.iter())
|
|
.map(|(f, s)| f - s)
|
|
.collect())
|
|
}
|
|
|
|
/// Calculate MACD signal line
|
|
fn calculate_macd_signal(
|
|
&self,
|
|
prices: &[f64],
|
|
fast: usize,
|
|
slow: usize,
|
|
signal: usize,
|
|
) -> Result<Vec<f64>> {
|
|
let macd = self.calculate_macd(prices, fast, slow)?;
|
|
self.calculate_ema(&macd, signal)
|
|
}
|
|
|
|
/// Calculate Bollinger upper band
|
|
fn calculate_bb_upper(&self, prices: &[f64], period: usize, num_std: f64) -> Result<Vec<f32>> {
|
|
let sma = self.calculate_sma(prices, period)?;
|
|
let mut upper = Vec::with_capacity(prices.len());
|
|
|
|
for i in 0..prices.len() {
|
|
if i < period - 1 {
|
|
upper.push(prices[i] as f32);
|
|
continue;
|
|
}
|
|
|
|
let window = &prices[i - period + 1..=i];
|
|
let mean = window.iter().sum::<f64>() / period as f64;
|
|
let variance = window.iter().map(|&x| (x - mean).powi(2)).sum::<f64>() / period as f64;
|
|
let std_dev = variance.sqrt();
|
|
|
|
upper.push((sma[i] + num_std * std_dev) as f32);
|
|
}
|
|
|
|
Ok(upper)
|
|
}
|
|
|
|
/// Calculate Bollinger lower band
|
|
fn calculate_bb_lower(&self, prices: &[f64], period: usize, num_std: f64) -> Result<Vec<f32>> {
|
|
let sma = self.calculate_sma(prices, period)?;
|
|
let mut lower = Vec::with_capacity(prices.len());
|
|
|
|
for i in 0..prices.len() {
|
|
if i < period - 1 {
|
|
lower.push(prices[i] as f32);
|
|
continue;
|
|
}
|
|
|
|
let window = &prices[i - period + 1..=i];
|
|
let mean = window.iter().sum::<f64>() / period as f64;
|
|
let variance = window.iter().map(|&x| (x - mean).powi(2)).sum::<f64>() / period as f64;
|
|
let std_dev = variance.sqrt();
|
|
|
|
lower.push((sma[i] - num_std * std_dev) as f32);
|
|
}
|
|
|
|
Ok(lower)
|
|
}
|
|
|
|
/// Calculate ATR (Average True Range)
|
|
fn calculate_atr(&self, bars: &[OHLCVBar], period: usize) -> Result<Vec<f32>> {
|
|
let mut atr = Vec::with_capacity(bars.len());
|
|
|
|
for i in 0..bars.len() {
|
|
if i < period {
|
|
atr.push(0.0);
|
|
continue;
|
|
}
|
|
|
|
let mut true_ranges = Vec::with_capacity(period);
|
|
for j in (i - period + 1)..=i {
|
|
let high_low = bars[j].high - bars[j].low;
|
|
let high_close = if j > 0 {
|
|
(bars[j].high - bars[j - 1].close).abs()
|
|
} else {
|
|
high_low
|
|
};
|
|
let low_close = if j > 0 {
|
|
(bars[j].low - bars[j - 1].close).abs()
|
|
} else {
|
|
high_low
|
|
};
|
|
|
|
let true_range = high_low.max(high_close).max(low_close);
|
|
true_ranges.push(true_range);
|
|
}
|
|
|
|
let avg_tr = true_ranges.iter().sum::<f64>() / period as f64;
|
|
atr.push(avg_tr as f32);
|
|
}
|
|
|
|
Ok(atr)
|
|
}
|
|
|
|
/// Calculate EMA (Exponential Moving Average)
|
|
///
|
|
/// Returns f64 to avoid precision-losing f64→f32→f64 round-trips
|
|
/// when the result feeds back into further f64 arithmetic (e.g. MACD).
|
|
fn calculate_ema(&self, prices: &[f64], period: usize) -> Result<Vec<f64>> {
|
|
if prices.len() < period {
|
|
return Err(anyhow::anyhow!(
|
|
"Need at least {} prices for EMA (got {})",
|
|
period,
|
|
prices.len()
|
|
));
|
|
}
|
|
|
|
let mut ema = Vec::with_capacity(prices.len());
|
|
let alpha = 2.0 / (period as f64 + 1.0);
|
|
|
|
// Start with SMA for first period
|
|
let initial_sma = prices[..period].iter().sum::<f64>() / period as f64;
|
|
ema.extend(vec![initial_sma; period]);
|
|
|
|
// Calculate EMA for remaining values
|
|
for i in period..prices.len() {
|
|
let new_ema = alpha * prices[i] + (1.0 - alpha) * ema[i - 1];
|
|
ema.push(new_ema);
|
|
}
|
|
|
|
Ok(ema)
|
|
}
|
|
|
|
/// Calculate SMA (Simple Moving Average)
|
|
///
|
|
/// Returns f64 to avoid precision-losing f64→f32→f64 round-trips
|
|
/// when the result feeds back into further f64 arithmetic (e.g. Bollinger bands).
|
|
fn calculate_sma(&self, prices: &[f64], period: usize) -> Result<Vec<f64>> {
|
|
let mut sma = Vec::with_capacity(prices.len());
|
|
|
|
for i in 0..prices.len() {
|
|
if i < period - 1 {
|
|
sma.push(prices[i]);
|
|
continue;
|
|
}
|
|
|
|
let window = &prices[i - period + 1..=i];
|
|
let avg = window.iter().sum::<f64>() / period as f64;
|
|
sma.push(avg);
|
|
}
|
|
|
|
Ok(sma)
|
|
}
|
|
|
|
// ===== Helper Methods =====
|
|
|
|
fn price_range(&self, bars: &[OHLCVBar]) -> (f64, f64) {
|
|
let mut min = f64::MAX;
|
|
let mut max = f64::MIN;
|
|
|
|
for bar in bars {
|
|
min = min.min(bar.low);
|
|
max = max.max(bar.high);
|
|
}
|
|
|
|
(min, max)
|
|
}
|
|
|
|
fn volume_range(&self, bars: &[OHLCVBar]) -> (f64, f64) {
|
|
let mut min = f64::MAX;
|
|
let mut max = f64::MIN;
|
|
|
|
for bar in bars {
|
|
min = min.min(bar.volume);
|
|
max = max.max(bar.volume);
|
|
}
|
|
|
|
(min, max)
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
|
|
/// ZN.FUT DBN data is not available locally; see `ml-asset-selection/src/lib.rs`
|
|
/// for the production universe definition that includes ZN.FUT.
|
|
#[tokio::test]
|
|
#[ignore = "ZN.FUT DBN data not available — see ml-asset-selection/src/lib.rs for the production universe definition"]
|
|
async fn test_load_symbol_data() -> Result<()> {
|
|
let mut loader = RealDataLoader::new_from_workspace()?;
|
|
|
|
// Test loading ZN.FUT (should have ~29K bars)
|
|
let bars = loader.load_symbol_data("ZN.FUT").await?;
|
|
assert!(bars.len() > 1000, "Expected >1000 bars, got {}", bars.len());
|
|
|
|
// Validate bar integrity
|
|
for bar in bars.iter().take(100) {
|
|
assert!(bar.high >= bar.low, "High < Low: {:?}", bar);
|
|
assert!(bar.high >= bar.open, "High < Open: {:?}", bar);
|
|
assert!(bar.high >= bar.close, "High < Close: {:?}", bar);
|
|
assert!(bar.low <= bar.open, "Low > Open: {:?}", bar);
|
|
assert!(bar.low <= bar.close, "Low > Close: {:?}", bar);
|
|
assert!(bar.volume >= 0.0, "Negative volume: {:?}", bar);
|
|
}
|
|
|
|
info!(bars = bars.len(), "Loaded bars for ZN.FUT");
|
|
Ok(())
|
|
}
|
|
|
|
/// ZN.FUT DBN data is not available locally; see `ml-asset-selection/src/lib.rs`
|
|
/// for the production universe definition that includes ZN.FUT.
|
|
#[tokio::test]
|
|
#[ignore = "ZN.FUT DBN data not available — see ml-asset-selection/src/lib.rs for the production universe definition"]
|
|
async fn test_extract_features() -> Result<()> {
|
|
let mut loader = RealDataLoader::new_from_workspace()?;
|
|
let bars = loader.load_symbol_data("ZN.FUT").await?;
|
|
|
|
let features = loader.extract_features(&bars)?;
|
|
|
|
assert_eq!(features.prices.len(), bars.len());
|
|
assert_eq!(features.returns.len(), bars.len());
|
|
assert_eq!(features.volume.len(), bars.len());
|
|
|
|
// Check normalization (should be 0-1 range)
|
|
for price_vec in features.prices.iter().take(100) {
|
|
for &val in price_vec {
|
|
assert!(val >= 0.0 && val <= 1.0, "Price not normalized: {}", val);
|
|
}
|
|
}
|
|
|
|
info!(bars = bars.len(), "Feature extraction working: 5 features/bar");
|
|
Ok(())
|
|
}
|
|
|
|
/// ZN.FUT DBN data is not available locally; see `ml-asset-selection/src/lib.rs`
|
|
/// for the production universe definition that includes ZN.FUT.
|
|
#[tokio::test]
|
|
#[ignore = "ZN.FUT DBN data not available — see ml-asset-selection/src/lib.rs for the production universe definition"]
|
|
async fn test_calculate_indicators() -> Result<()> {
|
|
let mut loader = RealDataLoader::new_from_workspace()?;
|
|
let bars = loader.load_symbol_data("ZN.FUT").await?;
|
|
|
|
let indicators = loader.calculate_indicators(&bars)?;
|
|
|
|
assert_eq!(indicators.rsi.len(), bars.len());
|
|
assert_eq!(indicators.macd.len(), bars.len());
|
|
assert_eq!(indicators.ema_fast.len(), bars.len());
|
|
|
|
// Check RSI validity (0-100 range)
|
|
for &rsi in indicators.rsi.iter().skip(14).take(100) {
|
|
assert!(rsi >= 0.0 && rsi <= 100.0, "Invalid RSI: {}", rsi);
|
|
}
|
|
|
|
// Check ATR validity (non-negative)
|
|
for &atr in indicators.atr.iter().skip(14).take(100) {
|
|
assert!(atr >= 0.0, "Invalid ATR: {}", atr);
|
|
}
|
|
|
|
info!(bars = bars.len(), "Indicators calculated: 10 indicators");
|
|
Ok(())
|
|
}
|
|
}
|