diff --git a/crates/fidc-core/src/data.rs b/crates/fidc-core/src/data.rs index ffa36ef..afa89f4 100644 --- a/crates/fidc-core/src/data.rs +++ b/crates/fidc-core/src/data.rs @@ -641,44 +641,40 @@ impl AdjustedCloseSeries { } fn current_moving_average(&self, date: NaiveDate, lookback: usize) -> Option { - if lookback == 0 { - return None; - } let end = match self.dates.binary_search(&date) { Ok(index) => index + 1, Err(0) => return None, Err(index) => index, }; - if end < lookback { - return None; - } - let base_factor = self.backward_factors.get(end - 1).copied().flatten()?; - let start = end - lookback; - if self.missing_back_adjusted_close_prefix[end] - != self.missing_back_adjusted_close_prefix[start] - { - return None; - } - let sum = self.back_adjusted_close_prefix[end] - self.back_adjusted_close_prefix[start]; - if !sum.is_finite() { - return None; - } - Some(normalize_rolling_factor( - sum / lookback as f64 / base_factor, - 12, - )) + self.moving_average_at_end(end, lookback) } fn decision_moving_average(&self, date: NaiveDate, lookback: usize) -> Option { - if lookback == 0 { - return None; - } let end = match self.dates.binary_search(&date) { Ok(index) => index, Err(0) => return None, Err(index) => index, }; - if end < lookback { + self.moving_average_at_end(end, lookback) + } + + fn moving_averages( + &self, + date: NaiveDate, + lookbacks: &[usize; N], + include_now: bool, + ) -> [Option; N] { + let end = match self.dates.binary_search(&date) { + Ok(index) if include_now => index + 1, + Ok(index) => index, + Err(0) => return [None; N], + Err(index) => index, + }; + std::array::from_fn(|index| self.moving_average_at_end(end, lookbacks[index])) + } + + fn moving_average_at_end(&self, end: usize, lookback: usize) -> Option { + if lookback == 0 || end < lookback { return None; } let base_factor = self.backward_factors.get(end - 1).copied().flatten()?; @@ -989,6 +985,27 @@ impl SymbolPriceSeries { }) } + fn volume_moving_averages( + &self, + date: NaiveDate, + lookbacks: &[usize; N], + include_now: bool, + ) -> [Option; N] { + let Some(end) = self.rolling_end_index(date, include_now) else { + return [None; N]; + }; + std::array::from_fn(|index| { + let lookback = lookbacks[index]; + self.valid_volume_window(end, lookback).map(|(start, end)| { + normalize_rolling_factor( + (self.valid_volume_sum_prefix[end] - self.valid_volume_sum_prefix[start]) + / lookback as f64, + 12, + ) + }) + }) + } + fn decision_volume_values(&self, date: NaiveDate, lookback: usize) -> Option> { let end = self.previous_completed_end_index(date)?; self.valid_volume_values(end, lookback) @@ -1031,6 +1048,15 @@ impl SymbolPriceSeries { } } + fn rolling_end_index(&self, date: NaiveDate, include_now: bool) -> Option { + match self.dates.binary_search(&date) { + Ok(index) if include_now => Some(index + 1), + Ok(index) => Some(index), + Err(0) => None, + Err(index) => Some(index), + } + } + fn price_values_for(&self, field: PriceField) -> &[f64] { match field { PriceField::DayOpen => &self.day_opens, @@ -1280,6 +1306,12 @@ pub(crate) struct SymbolSnapshotRefs<'a> { pub candidate: Option<&'a CandidateEligibility>, } +#[derive(Debug, Clone, Copy)] +pub(crate) struct StandardRollingMeans { + pub close: [Option; 7], + pub volume: [Option; 5], +} + impl DataSet { pub fn with_additional_trading_dates( mut self, @@ -1926,6 +1958,25 @@ impl DataSet { } } + pub(crate) fn market_standard_rolling_means_by_symbol_id( + &self, + date: NaiveDate, + symbol_id: u32, + close_lookbacks: &[usize; 7], + volume_lookbacks: &[usize; 5], + include_now: bool, + ) -> StandardRollingMeans { + let close = self + .adjusted_close_series_by_symbol_id(symbol_id) + .map(|series| series.moving_averages(date, close_lookbacks, include_now)) + .unwrap_or([None; 7]); + let volume = self + .market_series_by_symbol_id(symbol_id) + .map(|series| series.volume_moving_averages(date, volume_lookbacks, include_now)) + .unwrap_or([None; 5]); + StandardRollingMeans { close, volume } + } + pub fn benchmark(&self, date: NaiveDate) -> Option<&BenchmarkSnapshot> { self.benchmark_by_date.get(&date) } @@ -4960,6 +5011,121 @@ mod tests { .expect("volume contract dataset") } + #[test] + fn batched_standard_rolling_means_match_scalar_lookups() { + let dates = [ + NaiveDate::parse_from_str("2025-01-02", "%Y-%m-%d").unwrap(), + NaiveDate::parse_from_str("2025-01-03", "%Y-%m-%d").unwrap(), + NaiveDate::parse_from_str("2025-01-06", "%Y-%m-%d").unwrap(), + ]; + let data = DataSet::from_components( + vec![Instrument { + symbol: "000001.SZ".to_string(), + name: "000001.SZ".to_string(), + board: "SZ".to_string(), + round_lot: 100, + listed_at: Some(dates[0]), + delisted_at: None, + status: "active".to_string(), + }], + dates + .iter() + .enumerate() + .map(|(index, date)| { + market_row(&date.format("%Y-%m-%d").to_string(), 10.0 + index as f64, 100 + index as u64) + }) + .collect(), + dates + .iter() + .map(|date| DailyFactorSnapshot { + date: *date, + symbol: "000001.SZ".to_string(), + market_cap_bn: 10.0, + free_float_cap_bn: 8.0, + pe_ttm: 10.0, + turnover_ratio: None, + effective_turnover_ratio: None, + extra_factors: BTreeMap::from([(Cow::Borrowed(BACKWARD_ADJUSTMENT_FACTOR_FIELD), 1.0)]), + }) + .collect(), + Vec::new(), + dates + .iter() + .map(|date| BenchmarkSnapshot { + date: *date, + benchmark: "000852.SH".to_string(), + open: 100.0, + close: 100.0, + prev_close: 100.0, + volume: 1_000_000, + }) + .collect(), + ) + .expect("standard rolling dataset"); + let date = dates[2]; + let symbol_id = data.symbol_id("000001.SZ").unwrap(); + let close_lookbacks = [1, 2, 3, 1, 2, 3, 0]; + let volume_lookbacks = [1, 2, 3, 0, 2]; + let batched = data.market_standard_rolling_means_by_symbol_id( + date, + symbol_id, + &close_lookbacks, + &volume_lookbacks, + false, + ); + for (index, lookback) in close_lookbacks.iter().copied().enumerate() { + assert_eq!( + batched.close[index], + data.market_decision_numeric_moving_average_by_symbol_id( + date, + symbol_id, + "000001.SZ", + "close", + lookback, + ) + ); + } + for (index, lookback) in volume_lookbacks.iter().copied().enumerate() { + assert_eq!( + batched.volume[index], + data.market_decision_numeric_moving_average_by_symbol_id( + date, + symbol_id, + "000001.SZ", + "volume", + lookback, + ) + ); + } + let current = data.market_standard_rolling_means_by_symbol_id( + date, + symbol_id, + &close_lookbacks, + &volume_lookbacks, + true, + ); + assert_eq!( + current.close[1], + data.market_current_numeric_moving_average_by_symbol_id( + date, + symbol_id, + "000001.SZ", + "close", + 2, + ) + ); + assert_eq!( + current.volume[1], + data.market_current_numeric_moving_average_by_symbol_id( + date, + symbol_id, + "000001.SZ", + "volume", + 2, + ) + ); + } + #[test] fn source_volume_contract_rejects_windows_containing_missing_values() { let data = volume_contract_data(Some([1.0, 0.0, 1.0])); diff --git a/crates/fidc-core/src/platform_expr_strategy.rs b/crates/fidc-core/src/platform_expr_strategy.rs index 1774059..7eb9739 100644 --- a/crates/fidc-core/src/platform_expr_strategy.rs +++ b/crates/fidc-core/src/platform_expr_strategy.rs @@ -3937,57 +3937,49 @@ impl PlatformExprStrategy { None }; let instrument = ctx.data.instrument(symbol); - let rolling = |field: &'static str, lookback: usize| -> f64 { - if !self.stock_rolling_requirements.requires(field, lookback) { - return f64::NAN; - } - self.stock_decision_rolling_mean(ctx, date, symbol_id, symbol, field, lookback) - .unwrap_or(f64::NAN) + let required_rolling = |field: &'static str, lookback: usize| { + self.stock_rolling_requirements + .requires(field, lookback) + .then_some(lookback) + .unwrap_or(0) }; - let stock_ma_short = rolling("close", self.config.stock_short_ma_days); - let stock_ma_mid = rolling("close", self.config.stock_mid_ma_days); - let stock_ma_long = rolling("close", self.config.stock_long_ma_days); - let stock_ma5 = if self.config.stock_short_ma_days == 5 { - stock_ma_short - } else if self.config.stock_mid_ma_days == 5 { - stock_ma_mid - } else if self.config.stock_long_ma_days == 5 { - stock_ma_long - } else { - rolling("close", 5) - }; - let stock_ma10 = if self.config.stock_short_ma_days == 10 { - stock_ma_short - } else if self.config.stock_mid_ma_days == 10 { - stock_ma_mid - } else if self.config.stock_long_ma_days == 10 { - stock_ma_long - } else { - rolling("close", 10) - }; - let stock_ma20 = if self.config.stock_short_ma_days == 20 { - stock_ma_short - } else if self.config.stock_mid_ma_days == 20 { - stock_ma_mid - } else if self.config.stock_long_ma_days == 20 { - stock_ma_long - } else { - rolling("close", 20) - }; - let stock_ma30 = if self.config.stock_short_ma_days == 30 { - stock_ma_short - } else if self.config.stock_mid_ma_days == 30 { - stock_ma_mid - } else if self.config.stock_long_ma_days == 30 { - stock_ma_long - } else { - rolling("close", 30) - }; - let stock_volume_ma5 = rolling("volume", 5); - let stock_volume_ma10 = rolling("volume", 10); - let stock_volume_ma20 = rolling("volume", 20); - let stock_volume_ma60 = rolling("volume", 60); - let stock_volume_ma100 = rolling("volume", 100); + let close_lookbacks = [ + required_rolling("close", self.config.stock_short_ma_days), + required_rolling("close", self.config.stock_mid_ma_days), + required_rolling("close", self.config.stock_long_ma_days), + required_rolling("close", 5), + required_rolling("close", 10), + required_rolling("close", 20), + required_rolling("close", 30), + ]; + let volume_lookbacks = [ + required_rolling("volume", 5), + required_rolling("volume", 10), + required_rolling("volume", 20), + required_rolling("volume", 60), + required_rolling("volume", 100), + ]; + let rolling_means = ctx.data.market_standard_rolling_means_by_symbol_id( + date, + symbol_id, + &close_lookbacks, + &volume_lookbacks, + false, + ); + let close_rolling = |index: usize| rolling_means.close[index].unwrap_or(f64::NAN); + let volume_rolling = |index: usize| rolling_means.volume[index].unwrap_or(f64::NAN); + let stock_ma_short = close_rolling(0); + let stock_ma_mid = close_rolling(1); + let stock_ma_long = close_rolling(2); + let stock_ma5 = close_rolling(3); + let stock_ma10 = close_rolling(4); + let stock_ma20 = close_rolling(5); + let stock_ma30 = close_rolling(6); + let stock_volume_ma5 = volume_rolling(0); + let stock_volume_ma10 = volume_rolling(1); + let stock_volume_ma20 = volume_rolling(2); + let stock_volume_ma60 = volume_rolling(3); + let stock_volume_ma100 = volume_rolling(4); let touched_upper_limit = if intraday_same_day_factor { !market.paused && (market.is_at_upper_limit_price(market.close)