perf: scan signal rolling aggregates without allocations
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user