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()
|
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(
|
fn trailing_snapshots(
|
||||||
&self,
|
&self,
|
||||||
date: NaiveDate,
|
date: NaiveDate,
|
||||||
@@ -3606,6 +3649,32 @@ impl DataSet {
|
|||||||
.unwrap_or_default()
|
.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(
|
pub fn factor_numeric_values(
|
||||||
&self,
|
&self,
|
||||||
date: NaiveDate,
|
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]
|
#[test]
|
||||||
fn series_end_position_index_preserves_decision_and_current_boundaries() {
|
fn series_end_position_index_preserves_decision_and_current_boundaries() {
|
||||||
let data = volume_contract_data(Some([1.0, 1.0, 1.0]));
|
let data = volume_contract_data(Some([1.0, 1.0, 1.0]));
|
||||||
|
|||||||
@@ -5347,6 +5347,26 @@ impl PlatformExprStrategy {
|
|||||||
Ok(RuntimeHelperResolution::Number(value))
|
Ok(RuntimeHelperResolution::Number(value))
|
||||||
}
|
}
|
||||||
CompiledRuntimeHelperArgs::RollingMaxCurrent { field, lookback } => {
|
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 =
|
let values =
|
||||||
self.resolve_current_rolling_values(ctx, day, stock, field, *lookback)?;
|
self.resolve_current_rolling_values(ctx, day, stock, field, *lookback)?;
|
||||||
let value = values.iter().copied().fold(f64::NEG_INFINITY, f64::max);
|
let value = values.iter().copied().fold(f64::NEG_INFINITY, f64::max);
|
||||||
@@ -5356,6 +5376,35 @@ impl PlatformExprStrategy {
|
|||||||
field,
|
field,
|
||||||
return_count,
|
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(
|
let values = self.resolve_current_rolling_values(
|
||||||
ctx,
|
ctx,
|
||||||
day,
|
day,
|
||||||
|
|||||||
Reference in New Issue
Block a user