perf: scan signal rolling aggregates without allocations

This commit is contained in:
boris
2026-09-05 13:19:04 +08:00
parent f9ec86436a
commit b65b3ed8f1
2 changed files with 186 additions and 0 deletions
+137
View File
@@ -901,6 +901,49 @@ impl SymbolPriceSeries {
self.price_values_for(field)[start..end].to_vec()
}
fn close_max_at_end(&self, end: usize, lookback: usize) -> Option<f64> {
if lookback == 0 || end < lookback || end > self.closes.len() {
return None;
}
self.closes[end - lookback..end]
.iter()
.copied()
.fold(None, |maximum, value| {
Some(maximum.map_or(value, |current| current.max(value)))
})
}
fn close_return_sample_stddev_at_end(
&self,
end: usize,
lookback: usize,
) -> Result<Option<f64>, ()> {
if lookback == 0 || end < lookback || end > self.closes.len() {
return Ok(None);
}
if lookback < 2 {
return Ok(Some(0.0));
}
let values = &self.closes[end - lookback..end];
let mut sum = 0.0;
for pair in values.windows(2) {
let value = pair[1] / pair[0] - 1.0;
if !value.is_finite() {
return Err(());
}
sum += value;
}
let count = (lookback - 1) as f64;
let mean = sum / count;
let mut squared_sum = 0.0;
for pair in values.windows(2) {
let value = pair[1] / pair[0] - 1.0;
let difference = value - mean;
squared_sum += difference * difference;
}
Ok(Some((squared_sum / (count - 1.0)).sqrt()))
}
fn trailing_snapshots(
&self,
date: NaiveDate,
@@ -3606,6 +3649,32 @@ impl DataSet {
.unwrap_or_default()
}
pub(crate) fn market_current_close_max_by_symbol_id(
&self,
date: NaiveDate,
symbol_id: u32,
lookback: usize,
) -> Option<f64> {
let series = self.market_series_by_symbol_id(symbol_id)?;
let end = self.market_current_series_end_index_by_symbol_id(date, symbol_id)?;
series.close_max_at_end(end, lookback)
}
pub(crate) fn market_current_close_return_stddev_by_symbol_id(
&self,
date: NaiveDate,
symbol_id: u32,
lookback: usize,
) -> Result<Option<f64>, ()> {
let Some(series) = self.market_series_by_symbol_id(symbol_id) else {
return Ok(None);
};
let Some(end) = self.market_current_series_end_index_by_symbol_id(date, symbol_id) else {
return Ok(None);
};
series.close_return_sample_stddev_at_end(end, lookback)
}
pub fn factor_numeric_values(
&self,
date: NaiveDate,
@@ -6045,6 +6114,74 @@ mod tests {
);
}
#[test]
fn direct_current_close_aggregates_match_materialized_values() {
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(),
NaiveDate::parse_from_str("2025-01-07", "%Y-%m-%d").unwrap(),
];
let closes = [10.0, 11.0, 10.0, 12.0];
let market = dates
.iter()
.zip(closes)
.map(|(date, close)| market_row(&date.format("%Y-%m-%d").to_string(), close, 1_000))
.collect();
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(),
}],
market,
Vec::new(),
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("aggregate dataset");
let symbol_id = data.symbol_id("000001.SZ").expect("symbol id");
let materialized = data.market_closes_up_to(dates[3], "000001.SZ", 4);
assert_eq!(
data.market_current_close_max_by_symbol_id(dates[3], symbol_id, 4),
materialized.iter().copied().reduce(f64::max)
);
let returns = materialized
.windows(2)
.map(|pair| pair[1] / pair[0] - 1.0)
.collect::<Vec<_>>();
let mean = returns.iter().sum::<f64>() / returns.len() as f64;
let expected = (returns
.iter()
.map(|value| {
let difference = value - mean;
difference * difference
})
.sum::<f64>()
/ (returns.len() - 1) as f64)
.sqrt();
assert_eq!(
data.market_current_close_return_stddev_by_symbol_id(dates[3], symbol_id, 4)
.expect("valid returns"),
Some(expected)
);
}
#[test]
fn series_end_position_index_preserves_decision_and_current_boundaries() {
let data = volume_contract_data(Some([1.0, 1.0, 1.0]));
@@ -5347,6 +5347,26 @@ impl PlatformExprStrategy {
Ok(RuntimeHelperResolution::Number(value))
}
CompiledRuntimeHelperArgs::RollingMaxCurrent { field, lookback } => {
if field == "signal_close" {
let symbol_id =
ctx.data
.symbol_id(&self.config.signal_symbol)
.ok_or_else(|| {
BacktestError::Execution(format!(
"missing signal symbol {} for rolling max",
self.config.signal_symbol
))
})?;
let value = ctx
.data
.market_current_close_max_by_symbol_id(day.date, symbol_id, *lookback)
.ok_or_else(|| {
BacktestError::Execution(format!(
"missing current rolling values for field {field}"
))
})?;
return Ok(Self::normalized_runtime_number(value));
}
let values =
self.resolve_current_rolling_values(ctx, day, stock, field, *lookback)?;
let value = values.iter().copied().fold(f64::NEG_INFINITY, f64::max);
@@ -5356,6 +5376,35 @@ impl PlatformExprStrategy {
field,
return_count,
} => {
if field == "signal_close" {
let symbol_id =
ctx.data
.symbol_id(&self.config.signal_symbol)
.ok_or_else(|| {
BacktestError::Execution(format!(
"missing signal symbol {} for rolling return stddev",
self.config.signal_symbol
))
})?;
let value = ctx
.data
.market_current_close_return_stddev_by_symbol_id(
day.date,
symbol_id,
return_count.saturating_add(1),
)
.map_err(|_| {
BacktestError::Execution(format!(
"invalid current rolling return for field {field} with count {return_count}"
))
})?
.ok_or_else(|| {
BacktestError::Execution(format!(
"missing current rolling values for field {field}"
))
})?;
return Ok(Self::normalized_runtime_number(value));
}
let values = self.resolve_current_rolling_values(
ctx,
day,