perf(core): align market factor candidate lookups
This commit is contained in:
@@ -1239,6 +1239,13 @@ pub struct DataSet {
|
|||||||
futures_params_by_symbol: Arc<HashMap<String, Vec<FuturesTradingParameter>>>,
|
futures_params_by_symbol: Arc<HashMap<String, Vec<FuturesTradingParameter>>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Clone, Copy)]
|
||||||
|
pub(crate) struct SymbolSnapshotRefs<'a> {
|
||||||
|
pub market: Option<&'a DailyMarketSnapshot>,
|
||||||
|
pub factor: Option<&'a DailyFactorSnapshot>,
|
||||||
|
pub candidate: Option<&'a CandidateEligibility>,
|
||||||
|
}
|
||||||
|
|
||||||
impl DataSet {
|
impl DataSet {
|
||||||
pub fn with_additional_trading_dates(
|
pub fn with_additional_trading_dates(
|
||||||
mut self,
|
mut self,
|
||||||
@@ -1647,6 +1654,63 @@ impl DataSet {
|
|||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub(crate) fn symbol_snapshots_by_id(
|
||||||
|
&self,
|
||||||
|
date: NaiveDate,
|
||||||
|
symbol_id: u32,
|
||||||
|
) -> SymbolSnapshotRefs<'_> {
|
||||||
|
let market_rows = self.market_by_date.get(&date).map(Vec::as_slice);
|
||||||
|
let market_symbol_ids = self
|
||||||
|
.market_symbol_ids_by_date
|
||||||
|
.get(&date)
|
||||||
|
.map(Vec::as_slice);
|
||||||
|
let market_index = market_rows
|
||||||
|
.zip(market_symbol_ids)
|
||||||
|
.and_then(|(rows, symbol_ids)| symbol_id_index(rows.len(), symbol_ids, symbol_id));
|
||||||
|
let market = market_index.and_then(|index| market_rows?.get(index));
|
||||||
|
|
||||||
|
let factor = self
|
||||||
|
.factor_by_date
|
||||||
|
.get(&date)
|
||||||
|
.map(Vec::as_slice)
|
||||||
|
.zip(
|
||||||
|
self.factor_symbol_ids_by_date
|
||||||
|
.get(&date)
|
||||||
|
.map(Vec::as_slice),
|
||||||
|
)
|
||||||
|
.and_then(|(rows, symbol_ids)| {
|
||||||
|
find_by_symbol_id_with_preferred_index(
|
||||||
|
rows,
|
||||||
|
symbol_ids,
|
||||||
|
symbol_id,
|
||||||
|
market_index,
|
||||||
|
)
|
||||||
|
});
|
||||||
|
let candidate = self
|
||||||
|
.candidate_by_date
|
||||||
|
.get(&date)
|
||||||
|
.map(Vec::as_slice)
|
||||||
|
.zip(
|
||||||
|
self.candidate_symbol_ids_by_date
|
||||||
|
.get(&date)
|
||||||
|
.map(Vec::as_slice),
|
||||||
|
)
|
||||||
|
.and_then(|(rows, symbol_ids)| {
|
||||||
|
find_by_symbol_id_with_preferred_index(
|
||||||
|
rows,
|
||||||
|
symbol_ids,
|
||||||
|
symbol_id,
|
||||||
|
market_index,
|
||||||
|
)
|
||||||
|
});
|
||||||
|
|
||||||
|
SymbolSnapshotRefs {
|
||||||
|
market,
|
||||||
|
factor,
|
||||||
|
candidate,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
pub fn benchmark(&self, date: NaiveDate) -> Option<&BenchmarkSnapshot> {
|
pub fn benchmark(&self, date: NaiveDate) -> Option<&BenchmarkSnapshot> {
|
||||||
self.benchmark_by_date.get(&date)
|
self.benchmark_by_date.get(&date)
|
||||||
}
|
}
|
||||||
@@ -3554,9 +3618,30 @@ where
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn find_by_symbol_id<'a, T>(rows: &'a [T], symbol_ids: &[u32], symbol_id: u32) -> Option<&'a T> {
|
fn find_by_symbol_id<'a, T>(rows: &'a [T], symbol_ids: &[u32], symbol_id: u32) -> Option<&'a T> {
|
||||||
|
find_by_symbol_id_with_preferred_index(rows, symbol_ids, symbol_id, None)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn symbol_id_index(rows_len: usize, symbol_ids: &[u32], symbol_id: u32) -> Option<usize> {
|
||||||
|
if rows_len != symbol_ids.len() {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
symbol_ids.binary_search(&symbol_id).ok()
|
||||||
|
}
|
||||||
|
|
||||||
|
fn find_by_symbol_id_with_preferred_index<'a, T>(
|
||||||
|
rows: &'a [T],
|
||||||
|
symbol_ids: &[u32],
|
||||||
|
symbol_id: u32,
|
||||||
|
preferred_index: Option<usize>,
|
||||||
|
) -> Option<&'a T> {
|
||||||
if rows.len() != symbol_ids.len() {
|
if rows.len() != symbol_ids.len() {
|
||||||
return None;
|
return None;
|
||||||
}
|
}
|
||||||
|
if let Some(index) = preferred_index
|
||||||
|
&& symbol_ids.get(index).copied() == Some(symbol_id)
|
||||||
|
{
|
||||||
|
return rows.get(index);
|
||||||
|
}
|
||||||
symbol_ids
|
symbol_ids
|
||||||
.binary_search(&symbol_id)
|
.binary_search(&symbol_id)
|
||||||
.ok()
|
.ok()
|
||||||
@@ -3959,6 +4044,94 @@ mod tests {
|
|||||||
));
|
));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn combined_symbol_snapshot_lookup_uses_alignment_and_falls_back_for_sparse_rows() {
|
||||||
|
let date = NaiveDate::parse_from_str("2025-01-02", "%Y-%m-%d").unwrap();
|
||||||
|
let instrument = |symbol: &str| Instrument {
|
||||||
|
symbol: symbol.to_string(),
|
||||||
|
name: symbol.to_string(),
|
||||||
|
board: symbol
|
||||||
|
.rsplit_once('.')
|
||||||
|
.map(|(_, value)| value)
|
||||||
|
.unwrap_or("")
|
||||||
|
.to_string(),
|
||||||
|
round_lot: 100,
|
||||||
|
listed_at: None,
|
||||||
|
delisted_at: None,
|
||||||
|
status: "active".to_string(),
|
||||||
|
};
|
||||||
|
let market = |symbol: &str, close: f64| {
|
||||||
|
let mut row = market_row("2025-01-02", close, 1_000_000);
|
||||||
|
row.symbol = symbol.to_string();
|
||||||
|
row
|
||||||
|
};
|
||||||
|
let factor = |symbol: &str, market_cap_bn: f64| DailyFactorSnapshot {
|
||||||
|
date,
|
||||||
|
symbol: symbol.to_string(),
|
||||||
|
market_cap_bn,
|
||||||
|
free_float_cap_bn: market_cap_bn,
|
||||||
|
pe_ttm: 0.0,
|
||||||
|
turnover_ratio: None,
|
||||||
|
effective_turnover_ratio: None,
|
||||||
|
extra_factors: NumericFactorMap::new(),
|
||||||
|
};
|
||||||
|
let candidate = |symbol: &str| CandidateEligibility {
|
||||||
|
date,
|
||||||
|
symbol: symbol.to_string(),
|
||||||
|
is_st: false,
|
||||||
|
is_star_st: false,
|
||||||
|
is_new_listing: false,
|
||||||
|
is_paused: false,
|
||||||
|
allow_buy: true,
|
||||||
|
allow_sell: true,
|
||||||
|
is_kcb: false,
|
||||||
|
is_one_yuan: false,
|
||||||
|
risk_level_code: None,
|
||||||
|
};
|
||||||
|
let data = DataSet::from_components(
|
||||||
|
vec![
|
||||||
|
instrument("000001.SZ"),
|
||||||
|
instrument("000300.SH"),
|
||||||
|
instrument("600000.SH"),
|
||||||
|
],
|
||||||
|
vec![
|
||||||
|
market("000001.SZ", 10.0),
|
||||||
|
market("000300.SH", 20.0),
|
||||||
|
market("600000.SH", 12.0),
|
||||||
|
],
|
||||||
|
vec![factor("000001.SZ", 100.0), factor("600000.SH", 120.0)],
|
||||||
|
vec![candidate("000001.SZ"), candidate("600000.SH")],
|
||||||
|
vec![benchmark_row("2025-01-02", 20.0)],
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
|
for symbol in ["000001.SZ", "600000.SH"] {
|
||||||
|
let symbol_id = data.symbol_id(symbol).unwrap();
|
||||||
|
let combined = data.symbol_snapshots_by_id(date, symbol_id);
|
||||||
|
assert_eq!(
|
||||||
|
combined.market.map(|row| row.symbol.as_str()),
|
||||||
|
data.market_by_symbol_id(date, symbol_id)
|
||||||
|
.map(|row| row.symbol.as_str())
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
combined.factor.map(|row| row.symbol.as_str()),
|
||||||
|
data.factor_by_symbol_id(date, symbol_id)
|
||||||
|
.map(|row| row.symbol.as_str())
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
combined.candidate.map(|row| row.symbol.as_str()),
|
||||||
|
data.candidate_by_symbol_id(date, symbol_id)
|
||||||
|
.map(|row| row.symbol.as_str())
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
let signal_id = data.symbol_id("000300.SH").unwrap();
|
||||||
|
let signal = data.symbol_snapshots_by_id(date, signal_id);
|
||||||
|
assert_eq!(signal.market.map(|row| row.symbol.as_str()), Some("000300.SH"));
|
||||||
|
assert!(signal.factor.is_none());
|
||||||
|
assert!(signal.candidate.is_none());
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn additional_terminal_calendar_dates_are_isolated_from_shared_market_data() {
|
fn additional_terminal_calendar_dates_are_isolated_from_shared_market_data() {
|
||||||
let date = NaiveDate::parse_from_str("2025-01-02", "%Y-%m-%d").unwrap();
|
let date = NaiveDate::parse_from_str("2025-01-02", "%Y-%m-%d").unwrap();
|
||||||
|
|||||||
@@ -3898,13 +3898,34 @@ impl PlatformExprStrategy {
|
|||||||
return Ok(Arc::clone(state));
|
return Ok(Arc::clone(state));
|
||||||
}
|
}
|
||||||
|
|
||||||
let market = ctx
|
let execution_snapshots = ctx.data.symbol_snapshots_by_id(date, symbol_id);
|
||||||
.data
|
let market = execution_snapshots.market.ok_or_else(|| {
|
||||||
.require_market_by_symbol_id(date, symbol_id, symbol)?;
|
BacktestError::Data(crate::data::DataSetError::MissingSnapshot {
|
||||||
let feature_market = ctx
|
kind: "market",
|
||||||
.data
|
date,
|
||||||
.market_by_symbol_id(factor_date, symbol_id)
|
symbol: symbol.to_string(),
|
||||||
.unwrap_or(market);
|
})
|
||||||
|
})?;
|
||||||
|
let candidate = execution_snapshots.candidate.ok_or_else(|| {
|
||||||
|
BacktestError::Data(crate::data::DataSetError::MissingSnapshot {
|
||||||
|
kind: "candidate",
|
||||||
|
date,
|
||||||
|
symbol: symbol.to_string(),
|
||||||
|
})
|
||||||
|
})?;
|
||||||
|
let factor_snapshots = if factor_date == date {
|
||||||
|
execution_snapshots
|
||||||
|
} else {
|
||||||
|
ctx.data.symbol_snapshots_by_id(factor_date, symbol_id)
|
||||||
|
};
|
||||||
|
let feature_market = factor_snapshots.market.unwrap_or(market);
|
||||||
|
let factor = factor_snapshots.factor.ok_or_else(|| {
|
||||||
|
BacktestError::Data(crate::data::DataSetError::MissingSnapshot {
|
||||||
|
kind: "factor",
|
||||||
|
date: factor_date,
|
||||||
|
symbol: symbol.to_string(),
|
||||||
|
})
|
||||||
|
})?;
|
||||||
let intraday_same_day_factor = self.uses_intraday_execution_quotes()
|
let intraday_same_day_factor = self.uses_intraday_execution_quotes()
|
||||||
&& factor_date == date
|
&& factor_date == date
|
||||||
&& !ctx.is_lagged_execution();
|
&& !ctx.is_lagged_execution();
|
||||||
@@ -3913,12 +3934,6 @@ impl PlatformExprStrategy {
|
|||||||
} else {
|
} else {
|
||||||
None
|
None
|
||||||
};
|
};
|
||||||
let factor = ctx
|
|
||||||
.data
|
|
||||||
.require_factor_by_symbol_id(factor_date, symbol_id, symbol)?;
|
|
||||||
let candidate = ctx
|
|
||||||
.data
|
|
||||||
.require_candidate_by_symbol_id(date, symbol_id, symbol)?;
|
|
||||||
let instrument = ctx.data.instrument(symbol);
|
let instrument = ctx.data.instrument(symbol);
|
||||||
let rolling = |field: &'static str, lookback: usize| -> f64 {
|
let rolling = |field: &'static str, lookback: usize| -> f64 {
|
||||||
if !self.stock_rolling_requirements.requires(field, lookback) {
|
if !self.stock_rolling_requirements.requires(field, lookback) {
|
||||||
|
|||||||
Reference in New Issue
Block a user