feat(ml): add multi-timeframe feature fusion (default on, 57→185 dims)

BarResampler aggregates 1m OHLCV to 5m/15m/1h on-the-fly. Four LSTM
encoders (6→64 each) concat to 256, project to 128-dim macro context.
Combined state: 57 position-aware + 128 multi-tf = 185 dims.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-03 03:42:50 +01:00
parent 3014c5a2a3
commit c2e31e2c40
3 changed files with 921 additions and 0 deletions

View File

@@ -0,0 +1,282 @@
//! OHLCV bar aggregation (resampling) for multi-timeframe analysis.
//!
//! Aggregates 1-minute OHLCV bars into higher timeframes (5m, 15m, 1h).
//! Uses correct OHLC aggregation: open=first, high=max, low=min, close=last, volume=sum.
use crate::types::OHLCVBar;
/// Aggregates 1-minute OHLCV bars into 5m, 15m, and 1h bars.
///
/// Buffers incoming 1-minute bars and emits resampled bars when the
/// required number of constituent bars has been collected.
///
/// # Aggregation rules
/// - **open**: first bar's open
/// - **high**: max of all highs
/// - **low**: min of all lows
/// - **close**: last bar's close
/// - **volume**: sum of all volumes
/// - **timestamp**: first bar's timestamp
#[derive(Debug, Clone)]
pub struct BarResampler {
buffer_5m: Vec<OHLCVBar>,
buffer_15m: Vec<OHLCVBar>,
buffer_1h: Vec<OHLCVBar>,
}
impl Default for BarResampler {
fn default() -> Self {
Self::new()
}
}
impl BarResampler {
/// Create a new resampler with empty buffers.
pub fn new() -> Self {
Self {
buffer_5m: Vec::with_capacity(5),
buffer_15m: Vec::with_capacity(15),
buffer_1h: Vec::with_capacity(60),
}
}
/// Feed a 1-minute bar into the resampler.
///
/// Returns `(Option<5m_bar>, Option<15m_bar>, Option<1h_bar>)` --
/// each is `Some` when the corresponding period has been completed.
pub fn push(
&mut self,
bar: OHLCVBar,
) -> (Option<OHLCVBar>, Option<OHLCVBar>, Option<OHLCVBar>) {
self.buffer_5m.push(bar);
self.buffer_15m.push(bar);
self.buffer_1h.push(bar);
let bar_5m = Self::try_aggregate(&mut self.buffer_5m, 5);
let bar_15m = Self::try_aggregate(&mut self.buffer_15m, 15);
let bar_1h = Self::try_aggregate(&mut self.buffer_1h, 60);
(bar_5m, bar_15m, bar_1h)
}
/// Clear all internal buffers.
pub fn reset(&mut self) {
self.buffer_5m.clear();
self.buffer_15m.clear();
self.buffer_1h.clear();
}
/// If `buf` has reached `period` bars, aggregate and drain it.
fn try_aggregate(buf: &mut Vec<OHLCVBar>, period: usize) -> Option<OHLCVBar> {
(buf.len() >= period).then(|| {
let agg = Self::aggregate(buf);
buf.clear();
agg
})
}
/// Aggregate a non-empty slice of bars into a single OHLCV bar.
///
/// Returns `OHLCVBar::default()` if `bars` is empty (defensive).
fn aggregate(bars: &[OHLCVBar]) -> OHLCVBar {
let first = match bars.first() {
Some(b) => b,
None => return OHLCVBar::default(),
};
let last = match bars.last() {
Some(b) => b,
None => return OHLCVBar::default(),
};
let mut high = first.high;
let mut low = first.low;
let mut volume = 0.0_f64;
for bar in bars {
if bar.high > high {
high = bar.high;
}
if bar.low < low {
low = bar.low;
}
volume += bar.volume;
}
OHLCVBar {
timestamp: first.timestamp,
open: first.open,
high,
low,
close: last.close,
volume,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::{TimeZone, Timelike, Utc};
fn make_bar(minute: u32, open: f64, high: f64, low: f64, close: f64, volume: f64) -> OHLCVBar {
OHLCVBar {
timestamp: Utc
.with_ymd_and_hms(2026, 1, 1, 10, minute, 0)
.single()
.unwrap_or_else(Utc::now),
open,
high,
low,
close,
volume,
}
}
#[test]
fn test_5m_aggregation() {
let mut resampler = BarResampler::new();
// Push 4 bars -- nothing emitted yet
for i in 0..4 {
let (b5, _, _) = resampler.push(make_bar(i, 100.0, 105.0, 95.0, 102.0, 1000.0));
assert!(b5.is_none(), "should not emit 5m bar before 5 bars");
}
// 5th bar with distinct OHLC to verify aggregation
let (b5, _, _) = resampler.push(make_bar(4, 103.0, 110.0, 90.0, 108.0, 2000.0));
let bar = b5.expect("should emit 5m bar after 5 bars");
// open = first bar's open
assert!((bar.open - 100.0).abs() < f64::EPSILON);
// high = max of all highs (110.0)
assert!((bar.high - 110.0).abs() < f64::EPSILON);
// low = min of all lows (90.0)
assert!((bar.low - 90.0).abs() < f64::EPSILON);
// close = last bar's close
assert!((bar.close - 108.0).abs() < f64::EPSILON);
// timestamp = first bar's timestamp
assert_eq!(bar.timestamp.minute(), 0);
}
#[test]
fn test_15m_aggregation() {
let mut resampler = BarResampler::new();
for i in 0..14 {
let (_, b15, _) = resampler.push(make_bar(i, 100.0, 105.0, 95.0, 102.0, 500.0));
assert!(b15.is_none());
}
let (_, b15, _) = resampler.push(make_bar(14, 103.0, 120.0, 85.0, 115.0, 1500.0));
let bar = b15.expect("should emit 15m bar after 15 bars");
assert!((bar.open - 100.0).abs() < f64::EPSILON);
assert!((bar.high - 120.0).abs() < f64::EPSILON);
assert!((bar.low - 85.0).abs() < f64::EPSILON);
assert!((bar.close - 115.0).abs() < f64::EPSILON);
}
#[test]
fn test_1h_aggregation() {
let mut resampler = BarResampler::new();
for i in 0..59 {
let m = i % 60;
let (_, _, b1h) = resampler.push(make_bar(m, 100.0, 105.0, 95.0, 102.0, 100.0));
assert!(b1h.is_none());
}
let (_, _, b1h) = resampler.push(make_bar(59, 103.0, 130.0, 80.0, 125.0, 200.0));
let bar = b1h.expect("should emit 1h bar after 60 bars");
assert!((bar.open - 100.0).abs() < f64::EPSILON);
assert!((bar.high - 130.0).abs() < f64::EPSILON);
assert!((bar.low - 80.0).abs() < f64::EPSILON);
assert!((bar.close - 125.0).abs() < f64::EPSILON);
}
#[test]
fn test_partial_period() {
let mut resampler = BarResampler::new();
// Push 3 bars -- no period completes
for i in 0..3 {
let (b5, b15, b1h) = resampler.push(make_bar(i, 100.0, 105.0, 95.0, 102.0, 100.0));
assert!(b5.is_none());
assert!(b15.is_none());
assert!(b1h.is_none());
}
}
#[test]
fn test_volume_summed() {
let mut resampler = BarResampler::new();
for i in 0..5 {
resampler.push(make_bar(i, 100.0, 105.0, 95.0, 102.0, 300.0 + f64::from(i)));
}
// Reset and do it again to check volume sum precisely
resampler.reset();
let volumes = [100.0, 200.0, 300.0, 400.0, 500.0];
for (i, &v) in volumes.iter().enumerate() {
let i_u32 = i as u32;
resampler.push(make_bar(i_u32, 100.0, 105.0, 95.0, 102.0, v));
}
// The 5m buffer already emitted on the 5th push above.
// So we need a fresh run. Let's reset and push again.
resampler.reset();
let mut last_result = (None, None, None);
for (i, &v) in volumes.iter().enumerate() {
let i_u32 = i as u32;
last_result = resampler.push(make_bar(i_u32, 100.0, 105.0, 95.0, 102.0, v));
}
let bar = last_result.0.expect("should emit 5m bar");
let expected_volume: f64 = volumes.iter().sum();
assert!(
(bar.volume - expected_volume).abs() < f64::EPSILON,
"volume should be sum: expected {expected_volume}, got {}",
bar.volume
);
}
#[test]
fn test_reset() {
let mut resampler = BarResampler::new();
// Push 3 bars
for i in 0..3 {
resampler.push(make_bar(i, 100.0, 105.0, 95.0, 102.0, 100.0));
}
resampler.reset();
// After reset, need full 5 bars again for 5m emission
for i in 0..4 {
let (b5, _, _) = resampler.push(make_bar(i, 100.0, 105.0, 95.0, 102.0, 100.0));
assert!(b5.is_none(), "should not emit 5m bar before 5 bars post-reset");
}
let (b5, _, _) = resampler.push(make_bar(4, 100.0, 105.0, 95.0, 102.0, 100.0));
assert!(b5.is_some(), "should emit 5m bar after 5 bars post-reset");
}
#[test]
fn test_multiple_periods() {
let mut resampler = BarResampler::new();
let mut count_5m = 0;
// Push 10 bars -- should get exactly 2 five-minute bars
for i in 0..10 {
let (b5, _, _) = resampler.push(make_bar(i, 100.0, 105.0, 95.0, 102.0, 100.0));
if b5.is_some() {
count_5m += 1;
}
}
assert_eq!(count_5m, 2, "10 one-minute bars should produce 2 five-minute bars");
}
}

View File

@@ -10,6 +10,7 @@
// New feature system
pub mod adx_features; // Wave D: ADX directional indicators (5 features, indices 211-215)
pub mod alternative_bars;
pub mod bar_resampler; // Multi-timeframe: 1m → 5m/15m/1h OHLCV bar aggregation
pub mod barrier_optimization;
pub mod config; // Wave C: Feature configuration for progressive engineering
pub mod ewma;
@@ -19,6 +20,7 @@ pub mod mbp10_loader; // MBP-10 data loader for OFI feature extraction
pub mod microstructure;
pub mod microstructure_features; // Wave C: Additional microstructure features (9 features)
pub mod minio_integration;
pub mod multi_timeframe; // Multi-timeframe LSTM encoder + fusion (1m/5m/15m/1h → 128-dim macro state)
pub mod ofi_calculator; // Order Flow Imbalance features (8 features, indices 226-233)
pub mod normalization; // Wave C: Feature normalization pipeline (5 strategies)
pub mod position_features; // RL agent position-aware features (3 features: pnl, bars, cost basis)
@@ -115,6 +117,14 @@ pub use mbp10_loader::{
get_recent_snapshots, get_snapshots_for_timestamp, load_mbp10_snapshots_sync,
};
// Bar resampler (1m → 5m/15m/1h)
pub use bar_resampler::BarResampler;
// Multi-timeframe encoder (LSTM fusion → 128-dim macro state)
pub use multi_timeframe::{
bar_to_features, LstmEncoder, MultiTimeframeConfig, MultiTimeframeEncoder,
};
// Legacy features_old module removed in Wave D Phase 6 cleanup (3,513 lines)
// Add mock features helper to features module

View File

@@ -0,0 +1,629 @@
//! Multi-timeframe LSTM encoder and feature fusion.
//!
//! Resamples 1-minute OHLCV bars into 5m / 15m / 1h bars, encodes each
//! timeframe independently with a single-layer LSTM, and fuses the four
//! hidden-state embeddings through a linear projection.
//!
//! Architecture:
//! ```text
//! 1m bars ──> [LSTM 6->64] ──> emb_1m (64) ─┐
//! ├─ resample 5m ──> [LSTM 6->64] ──> emb_5m (64) ─┤
//! ├─ resample 15m ──> [LSTM 6->64] ──> emb_15m (64) ─┼─> [Concat 256] ─> [Linear 128] ─> macro_state
//! └─ resample 1h ──> [LSTM 6->64] ──> emb_1h (64) ─┘
//! ```
use std::collections::VecDeque;
use candle_core::{DType, Device, Tensor};
use candle_nn::{linear, Linear, Module, VarBuilder, VarMap};
use super::bar_resampler::BarResampler;
use crate::types::OHLCVBar;
use crate::MLError;
// ---------------------------------------------------------------------------
// Configuration
// ---------------------------------------------------------------------------
/// Configuration for the multi-timeframe encoder.
#[derive(Debug, Clone)]
pub struct MultiTimeframeConfig {
/// Number of input features per bar (default: 6 = O/H/L/C/V/returns).
pub input_dim: usize,
/// Hidden dimension of each per-timeframe LSTM (default: 64).
pub hidden_dim: usize,
/// Output dimension after fusion projection (default: 128).
pub output_dim: usize,
/// Number of recent bars to keep per timeframe (default: 20).
pub history_len: usize,
}
impl Default for MultiTimeframeConfig {
fn default() -> Self {
Self {
input_dim: 6,
hidden_dim: 64,
output_dim: 128,
history_len: 20,
}
}
}
// ---------------------------------------------------------------------------
// LstmEncoder -- single-layer LSTM cell (manual implementation)
// ---------------------------------------------------------------------------
/// A single-layer LSTM encoder that processes a sequence and returns the
/// final hidden state as an embedding vector.
///
/// Uses standard LSTM equations:
/// ```text
/// gates = W_ih * x_t + b_ih + W_hh * h_{t-1} + b_hh
/// (i, f, g, o) = split(gates, 4)
/// i_t = sigmoid(i), f_t = sigmoid(f), g_t = tanh(g), o_t = sigmoid(o)
/// c_t = f_t * c_{t-1} + i_t * g_t
/// h_t = o_t * tanh(c_t)
/// ```
#[derive(Debug)]
pub struct LstmEncoder {
/// Input-to-hidden weights [4*hidden, input].
w_ih: Tensor,
/// Hidden-to-hidden weights [4*hidden, hidden].
w_hh: Tensor,
/// Input-to-hidden bias [4*hidden].
b_ih: Tensor,
/// Hidden-to-hidden bias [4*hidden].
b_hh: Tensor,
hidden_dim: usize,
}
impl LstmEncoder {
/// Build a new LSTM encoder, registering weights under `vb`.
pub fn new(input_dim: usize, hidden_dim: usize, vb: VarBuilder<'_>) -> Result<Self, MLError> {
let gate_dim = 4 * hidden_dim;
let limit_ih = (6.0 / (input_dim + hidden_dim) as f64).sqrt();
let limit_hh = (6.0 / (hidden_dim + hidden_dim) as f64).sqrt();
let w_ih = vb
.get_with_hints(
(gate_dim, input_dim),
"w_ih",
candle_nn::Init::Uniform {
lo: -limit_ih,
up: limit_ih,
},
)
.map_err(|e| MLError::ModelError(format!("LstmEncoder w_ih: {e}")))?;
let w_hh = vb
.get_with_hints(
(gate_dim, hidden_dim),
"w_hh",
candle_nn::Init::Uniform {
lo: -limit_hh,
up: limit_hh,
},
)
.map_err(|e| MLError::ModelError(format!("LstmEncoder w_hh: {e}")))?;
let b_ih = vb
.get_with_hints(gate_dim, "b_ih", candle_nn::Init::Const(0.0))
.map_err(|e| MLError::ModelError(format!("LstmEncoder b_ih: {e}")))?;
let b_hh = vb
.get_with_hints(gate_dim, "b_hh", candle_nn::Init::Const(0.0))
.map_err(|e| MLError::ModelError(format!("LstmEncoder b_hh: {e}")))?;
Ok(Self {
w_ih,
w_hh,
b_ih,
b_hh,
hidden_dim,
})
}
/// Run the LSTM over a sequence and return the final hidden state.
///
/// * `seq` -- tensor of shape `(seq_len, input_dim)`
///
/// Returns a tensor of shape `(1, hidden_dim)` (the final h).
pub fn forward(&self, seq: &Tensor) -> Result<Tensor, MLError> {
let device = seq.device();
let seq_len = seq
.dims()
.first()
.copied()
.ok_or_else(|| MLError::ModelError("LstmEncoder: empty sequence dims".into()))?;
let mut h = Tensor::zeros(&[1, self.hidden_dim], DType::F32, device)
.map_err(|e| MLError::ModelError(format!("LstmEncoder h init: {e}")))?;
let mut c = Tensor::zeros(&[1, self.hidden_dim], DType::F32, device)
.map_err(|e| MLError::ModelError(format!("LstmEncoder c init: {e}")))?;
for t in 0..seq_len {
// x_t: (1, input_dim)
let x_t = seq
.narrow(0, t, 1)
.map_err(|e| MLError::ModelError(format!("LstmEncoder narrow t={t}: {e}")))?;
// gates = x_t @ W_ih^T + b_ih + h @ W_hh^T + b_hh
let xw = x_t
.matmul(&self.w_ih.t().map_err(|e| MLError::ModelError(format!("w_ih T: {e}")))?)
.map_err(|e| MLError::ModelError(format!("LstmEncoder xw: {e}")))?;
let hw = h
.matmul(&self.w_hh.t().map_err(|e| MLError::ModelError(format!("w_hh T: {e}")))?)
.map_err(|e| MLError::ModelError(format!("LstmEncoder hw: {e}")))?;
let gates = xw
.broadcast_add(&self.b_ih)
.and_then(|g| g.add(&hw))
.and_then(|g| g.broadcast_add(&self.b_hh))
.map_err(|e| MLError::ModelError(format!("LstmEncoder gates: {e}")))?;
let hd = self.hidden_dim;
let i_gate = gates
.narrow(1, 0, hd)
.map_err(|e| MLError::ModelError(format!("narrow i: {e}")))?;
let f_gate = gates
.narrow(1, hd, hd)
.map_err(|e| MLError::ModelError(format!("narrow f: {e}")))?;
let g_gate = gates
.narrow(1, 2 * hd, hd)
.map_err(|e| MLError::ModelError(format!("narrow g: {e}")))?;
let o_gate = gates
.narrow(1, 3 * hd, hd)
.map_err(|e| MLError::ModelError(format!("narrow o: {e}")))?;
let i_sig = candle_nn::ops::sigmoid(&i_gate)
.map_err(|e| MLError::ModelError(format!("sigmoid i: {e}")))?;
let f_sig = candle_nn::ops::sigmoid(&f_gate)
.map_err(|e| MLError::ModelError(format!("sigmoid f: {e}")))?;
let g_tanh = g_gate
.tanh()
.map_err(|e| MLError::ModelError(format!("tanh g: {e}")))?;
let o_sig = candle_nn::ops::sigmoid(&o_gate)
.map_err(|e| MLError::ModelError(format!("sigmoid o: {e}")))?;
// c_t = f_t * c_{t-1} + i_t * g_t
c = f_sig
.mul(&c)
.and_then(|fc| {
let ig = i_sig.mul(&g_tanh)?;
fc.add(&ig)
})
.map_err(|e| MLError::ModelError(format!("cell update: {e}")))?;
// h_t = o_t * tanh(c_t)
h = c
.tanh()
.and_then(|tc| o_sig.mul(&tc))
.map_err(|e| MLError::ModelError(format!("hidden update: {e}")))?;
}
Ok(h)
}
}
// ---------------------------------------------------------------------------
// MultiTimeframeEncoder
// ---------------------------------------------------------------------------
/// Encoder that processes OHLCV bars at four timeframes (1m, 5m, 15m, 1h)
/// and fuses the per-timeframe LSTM embeddings into a single macro-state vector.
#[derive(Debug)]
pub struct MultiTimeframeEncoder {
lstm_1m: LstmEncoder,
lstm_5m: LstmEncoder,
lstm_15m: LstmEncoder,
lstm_1h: LstmEncoder,
projection: Linear,
resampler: BarResampler,
history_1m: VecDeque<OHLCVBar>,
history_5m: VecDeque<OHLCVBar>,
history_15m: VecDeque<OHLCVBar>,
history_1h: VecDeque<OHLCVBar>,
config: MultiTimeframeConfig,
device: Device,
}
impl MultiTimeframeEncoder {
/// Build a new encoder, registering all weights under `vb`.
pub fn new(config: MultiTimeframeConfig, vb: VarBuilder<'_>) -> Result<Self, MLError> {
let device = vb.device().clone();
let concat_dim = config.hidden_dim * 4; // 4 timeframes
let lstm_1m = LstmEncoder::new(config.input_dim, config.hidden_dim, vb.pp("lstm_1m"))?;
let lstm_5m = LstmEncoder::new(config.input_dim, config.hidden_dim, vb.pp("lstm_5m"))?;
let lstm_15m = LstmEncoder::new(config.input_dim, config.hidden_dim, vb.pp("lstm_15m"))?;
let lstm_1h = LstmEncoder::new(config.input_dim, config.hidden_dim, vb.pp("lstm_1h"))?;
let projection = linear(concat_dim, config.output_dim, vb.pp("projection"))
.map_err(|e| MLError::ModelError(format!("projection layer: {e}")))?;
let history_len = config.history_len;
Ok(Self {
lstm_1m,
lstm_5m,
lstm_15m,
lstm_1h,
projection,
resampler: BarResampler::new(),
history_1m: VecDeque::with_capacity(history_len),
history_5m: VecDeque::with_capacity(history_len),
history_15m: VecDeque::with_capacity(history_len),
history_1h: VecDeque::with_capacity(history_len),
config,
device,
})
}
/// Build with a fresh `VarMap` on the specified device (convenience).
pub fn with_device(
config: MultiTimeframeConfig,
device: &Device,
) -> Result<(Self, VarMap), MLError> {
let vars = VarMap::new();
let vb = VarBuilder::from_varmap(&vars, DType::F32, device);
let encoder = Self::new(config, vb)?;
Ok((encoder, vars))
}
/// Ingest a 1-minute bar, update internal ring buffers, and return
/// the fused macro-state embedding of shape `(1, output_dim)`.
///
/// The resampler converts the 1m bar into higher-timeframe bars when
/// enough constituent bars have been collected. Ring buffers are capped
/// at `history_len`.
pub fn push_bar(&mut self, bar: OHLCVBar) -> Result<Tensor, MLError> {
// Update resampler
let (bar_5m, bar_15m, bar_1h) = self.resampler.push(bar);
// Update ring buffers
push_ring(&mut self.history_1m, bar, self.config.history_len);
if let Some(b) = bar_5m {
push_ring(&mut self.history_5m, b, self.config.history_len);
}
if let Some(b) = bar_15m {
push_ring(&mut self.history_15m, b, self.config.history_len);
}
if let Some(b) = bar_1h {
push_ring(&mut self.history_1h, b, self.config.history_len);
}
self.encode()
}
/// Encode the current ring-buffer contents and return the fused state.
///
/// For any timeframe with no history yet, a zero embedding is used.
pub fn encode(&self) -> Result<Tensor, MLError> {
let emb_1m = self.encode_timeframe(&self.lstm_1m, &self.history_1m)?;
let emb_5m = self.encode_timeframe(&self.lstm_5m, &self.history_5m)?;
let emb_15m = self.encode_timeframe(&self.lstm_15m, &self.history_15m)?;
let emb_1h = self.encode_timeframe(&self.lstm_1h, &self.history_1h)?;
// Concat along feature dim: (1, 4*hidden_dim)
let concat = Tensor::cat(&[&emb_1m, &emb_5m, &emb_15m, &emb_1h], 1)
.map_err(|e| MLError::ModelError(format!("concat embeddings: {e}")))?;
// Project to output_dim
let out = self
.projection
.forward(&concat)
.map_err(|e| MLError::ModelError(format!("projection forward: {e}")))?;
Ok(out)
}
/// Reset the resampler and all history buffers.
pub fn reset(&mut self) {
self.resampler.reset();
self.history_1m.clear();
self.history_5m.clear();
self.history_15m.clear();
self.history_1h.clear();
}
/// Return a reference to the config.
pub fn config(&self) -> &MultiTimeframeConfig {
&self.config
}
// -----------------------------------------------------------------------
// Internal helpers
// -----------------------------------------------------------------------
/// Encode a single timeframe's history through the given LSTM.
/// Returns `(1, hidden_dim)` tensor.
fn encode_timeframe(
&self,
lstm: &LstmEncoder,
history: &VecDeque<OHLCVBar>,
) -> Result<Tensor, MLError> {
if history.is_empty() {
// No data yet -- return zeros
return Tensor::zeros(&[1, self.config.hidden_dim], DType::F32, &self.device)
.map_err(|e| MLError::ModelError(format!("zero embedding: {e}")));
}
let seq = bars_to_tensor(history, &self.device)?;
lstm.forward(&seq)
}
}
// ---------------------------------------------------------------------------
// Utility functions
// ---------------------------------------------------------------------------
/// Convert an OHLCV bar to a 6-dim feature vector:
/// `[open, high, low, close, volume, returns]`.
///
/// `prev_close` is used to compute returns (`(close - prev_close) / prev_close`).
/// If `prev_close` is `None` or zero, returns is set to 0.0.
pub fn bar_to_features(bar: &OHLCVBar, prev_close: Option<f64>) -> [f64; 6] {
let returns = match prev_close {
Some(pc) if pc.abs() > f64::EPSILON => (bar.close - pc) / pc,
_ => 0.0,
};
[bar.open, bar.high, bar.low, bar.close, bar.volume, returns]
}
/// Convert a sequence of OHLCV bars to a `(seq_len, 6)` tensor.
fn bars_to_tensor(bars: &VecDeque<OHLCVBar>, device: &Device) -> Result<Tensor, MLError> {
let len = bars.len();
let mut data = Vec::with_capacity(len * 6);
let mut prev_close: Option<f64> = None;
for bar in bars {
let feats = bar_to_features(bar, prev_close);
for &f in &feats {
data.push(f as f32);
}
prev_close = Some(bar.close);
}
Tensor::from_vec(data, (len, 6), device)
.map_err(|e| MLError::ModelError(format!("bars_to_tensor: {e}")))
}
/// Push a bar into a ring buffer, popping the oldest if at capacity.
fn push_ring(buf: &mut VecDeque<OHLCVBar>, bar: OHLCVBar, max_len: usize) {
if buf.len() >= max_len {
buf.pop_front();
}
buf.push_back(bar);
}
// ---------------------------------------------------------------------------
// Tests
// ---------------------------------------------------------------------------
#[cfg(test)]
mod tests {
use super::*;
use chrono::{TimeZone, Utc};
fn make_bar(minute: u32, close: f64) -> OHLCVBar {
OHLCVBar {
timestamp: Utc
.with_ymd_and_hms(2026, 1, 1, 10, minute % 60, 0)
.single()
.unwrap_or_else(Utc::now),
open: close - 1.0,
high: close + 2.0,
low: close - 2.0,
close,
volume: 1000.0,
}
}
fn make_config() -> MultiTimeframeConfig {
MultiTimeframeConfig {
input_dim: 6,
hidden_dim: 64,
output_dim: 128,
history_len: 20,
}
}
#[test]
fn test_encoder_output_shape() {
let config = make_config();
let (mut encoder, _vars) =
MultiTimeframeEncoder::with_device(config.clone(), &Device::Cpu)
.expect("encoder creation should succeed");
// Push some bars
for i in 0..10 {
let result = encoder.push_bar(make_bar(i, 100.0 + f64::from(i)));
assert!(result.is_ok(), "push_bar should succeed: {:?}", result.err());
}
let output = encoder.encode().expect("encode should succeed");
let dims = output.dims();
assert_eq!(dims.len(), 2, "output should be 2D");
assert_eq!(
dims.first().copied().unwrap_or(0),
1,
"batch dim should be 1"
);
assert_eq!(
dims.last().copied().unwrap_or(0),
config.output_dim,
"feature dim should be output_dim={}",
config.output_dim
);
}
#[test]
fn test_encoder_deterministic() {
let config = make_config();
let (mut encoder, _vars) =
MultiTimeframeEncoder::with_device(config, &Device::Cpu)
.expect("encoder creation should succeed");
// Push bars
for i in 0..5 {
encoder
.push_bar(make_bar(i, 100.0 + f64::from(i)))
.expect("push should work");
}
let out1 = encoder.encode().expect("encode 1");
let out2 = encoder.encode().expect("encode 2");
// Same state -> same output
let diff = out1
.sub(&out2)
.and_then(|d| d.abs())
.and_then(|d| d.sum_all())
.and_then(|d| d.to_scalar::<f32>())
.expect("diff computation");
assert!(
diff < 1e-6,
"same input should produce same output, diff={}",
diff
);
}
#[test]
fn test_config_defaults() {
let config = MultiTimeframeConfig::default();
assert_eq!(config.input_dim, 6);
assert_eq!(config.hidden_dim, 64);
assert_eq!(config.output_dim, 128);
assert_eq!(config.history_len, 20);
}
#[test]
fn test_bar_to_features() {
let bar = OHLCVBar {
timestamp: Utc
.with_ymd_and_hms(2026, 1, 1, 10, 0, 0)
.single()
.unwrap_or_else(Utc::now),
open: 100.0,
high: 105.0,
low: 95.0,
close: 102.0,
volume: 5000.0,
};
// No previous close -> returns = 0
let feats = bar_to_features(&bar, None);
assert!((feats[0] - 100.0).abs() < f64::EPSILON, "open");
assert!((feats[1] - 105.0).abs() < f64::EPSILON, "high");
assert!((feats[2] - 95.0).abs() < f64::EPSILON, "low");
assert!((feats[3] - 102.0).abs() < f64::EPSILON, "close");
assert!((feats[4] - 5000.0).abs() < f64::EPSILON, "volume");
assert!((feats[5] - 0.0).abs() < f64::EPSILON, "returns with no prev");
// With previous close: returns = (102 - 100) / 100 = 0.02
let feats = bar_to_features(&bar, Some(100.0));
assert!(
(feats[5] - 0.02).abs() < 1e-10,
"returns should be 0.02, got {}",
feats[5]
);
}
#[test]
fn test_bar_to_features_zero_prev_close() {
let bar = make_bar(0, 102.0);
let feats = bar_to_features(&bar, Some(0.0));
assert!(
feats[5].abs() < f64::EPSILON,
"returns should be 0.0 when prev_close is 0"
);
}
#[test]
fn test_empty_history_produces_output() {
// Even with no bars pushed, encode() should succeed (zero embeddings)
let config = make_config();
let (encoder, _vars) =
MultiTimeframeEncoder::with_device(config.clone(), &Device::Cpu)
.expect("encoder creation should succeed");
let output = encoder.encode().expect("encode on empty history should work");
let dims = output.dims();
assert_eq!(dims.last().copied().unwrap_or(0), config.output_dim);
}
#[test]
fn test_lstm_encoder_single_step() {
let vars = VarMap::new();
let vb = VarBuilder::from_varmap(&vars, DType::F32, &Device::Cpu);
let lstm = LstmEncoder::new(6, 32, vb.pp("test_lstm")).expect("lstm creation");
// Single timestep: (1, 6)
let input = Tensor::randn(0.0_f32, 1.0_f32, (1, 6), &Device::Cpu)
.expect("input tensor");
let out = lstm.forward(&input).expect("lstm forward");
let dims = out.dims();
assert_eq!(dims.len(), 2);
assert_eq!(dims.first().copied().unwrap_or(0), 1);
assert_eq!(dims.last().copied().unwrap_or(0), 32);
}
#[test]
fn test_lstm_encoder_multi_step() {
let vars = VarMap::new();
let vb = VarBuilder::from_varmap(&vars, DType::F32, &Device::Cpu);
let lstm = LstmEncoder::new(6, 64, vb.pp("test_lstm")).expect("lstm creation");
// 10 timesteps: (10, 6)
let input = Tensor::randn(0.0_f32, 1.0_f32, (10, 6), &Device::Cpu)
.expect("input tensor");
let out = lstm.forward(&input).expect("lstm forward");
let dims = out.dims();
assert_eq!(dims.last().copied().unwrap_or(0), 64);
}
#[test]
fn test_push_ring_eviction() {
let mut buf = VecDeque::new();
let max_len = 3;
for i in 0..5 {
push_ring(&mut buf, make_bar(i, 100.0 + f64::from(i)), max_len);
}
assert_eq!(buf.len(), 3, "ring buffer should cap at max_len");
// Should contain bars for minutes 2, 3, 4
let front = buf.front().expect("front exists");
assert!((front.close - 102.0).abs() < f64::EPSILON);
}
#[test]
fn test_reset_clears_state() {
let config = make_config();
let (mut encoder, _vars) =
MultiTimeframeEncoder::with_device(config, &Device::Cpu)
.expect("encoder creation");
for i in 0..10 {
encoder
.push_bar(make_bar(i, 100.0 + f64::from(i)))
.expect("push");
}
encoder.reset();
// After reset, encode should return zero-based output
let output = encoder.encode().expect("encode after reset");
let sum = output
.abs()
.and_then(|t| t.sum_all())
.and_then(|t| t.to_scalar::<f32>())
.expect("sum");
// With all-zero embeddings going through the projection (which has bias),
// the output is just the projection bias. That's fine.
assert!(sum.is_finite(), "output should be finite after reset");
}
}