线性构建数据集价格序列

This commit is contained in:
boris
2026-08-25 21:55:18 +08:00
parent 01cffb947c
commit ac30d86b6a
+127 -46
View File
@@ -14,6 +14,8 @@ use crate::futures::FuturesTradingParameter;
use crate::instrument::Instrument; use crate::instrument::Instrument;
use crate::risk_control::{ChinaAShareRiskControl, FidcRiskControlConfig}; use crate::risk_control::{ChinaAShareRiskControl, FidcRiskControlConfig};
const BACKWARD_ADJUSTMENT_FACTOR_FIELD: &str = "adjustment_factor_backward1";
mod date_format { mod date_format {
use chrono::NaiveDate; use chrono::NaiveDate;
use serde::{self, Deserialize, Deserializer, Serializer}; use serde::{self, Deserialize, Deserializer, Serializer};
@@ -504,21 +506,35 @@ struct AdjustedCloseSeries {
} }
impl AdjustedCloseSeries { impl AdjustedCloseSeries {
fn new( fn new(market: &SymbolPriceSeries, factor_rows: &[&DailyFactorSnapshot]) -> Option<Self> {
market: &SymbolPriceSeries, debug_assert!(
factor_by_date: &BTreeMap<NaiveDate, Vec<DailyFactorSnapshot>>, factor_rows
) -> Option<Self> { .windows(2)
.all(|window| window[0].date <= window[1].date)
);
let mut backward_factors = Vec::with_capacity(market.dates.len()); let mut backward_factors = Vec::with_capacity(market.dates.len());
let mut back_adjusted_closes = Vec::with_capacity(market.dates.len()); let mut back_adjusted_closes = Vec::with_capacity(market.dates.len());
let mut back_adjusted_close_prefix = Vec::with_capacity(market.dates.len() + 1); let mut back_adjusted_close_prefix = Vec::with_capacity(market.dates.len() + 1);
let mut missing_back_adjusted_close_prefix = Vec::with_capacity(market.dates.len() + 1); let mut missing_back_adjusted_close_prefix = Vec::with_capacity(market.dates.len() + 1);
back_adjusted_close_prefix.push(0.0); back_adjusted_close_prefix.push(0.0);
missing_back_adjusted_close_prefix.push(0); missing_back_adjusted_close_prefix.push(0);
let mut factor_index = 0usize;
for (date, close) in market.dates.iter().zip(&market.closes) { for (date, close) in market.dates.iter().zip(&market.closes) {
let factor = factor_by_date while factor_rows
.get(date) .get(factor_index)
.and_then(|rows| find_by_symbol(rows, &market.symbol, |row| row.symbol.as_str())) .is_some_and(|snapshot| snapshot.date < *date)
.and_then(|snapshot| factor_numeric_value(snapshot, "adjustment_factor_backward1")) {
factor_index += 1;
}
let factor = factor_rows
.get(factor_index)
.filter(|snapshot| snapshot.date == *date)
.and_then(|snapshot| {
snapshot
.extra_factors
.get(BACKWARD_ADJUSTMENT_FACTOR_FIELD)
.copied()
})
.filter(|factor| factor.is_finite() && *factor > 0.0); .filter(|factor| factor.is_finite() && *factor > 0.0);
let back_adjusted_close = factor let back_adjusted_close = factor
.filter(|_| close.is_finite() && *close > 0.0) .filter(|_| close.is_finite() && *close > 0.0)
@@ -651,42 +667,64 @@ impl AdjustedCloseSeries {
} }
impl SymbolPriceSeries { impl SymbolPriceSeries {
#[cfg(test)]
fn new<'a, I>(symbol: String, rows: I) -> Self fn new<'a, I>(symbol: String, rows: I) -> Self
where where
I: IntoIterator<Item = &'a DailyMarketSnapshot>, I: IntoIterator<Item = &'a DailyMarketSnapshot>,
{ {
let mut sorted = rows.into_iter().collect::<Vec<_>>(); let mut sorted = rows.into_iter().collect::<Vec<_>>();
sorted.sort_by_key(|row| row.date); sorted.sort_by_key(|row| row.date);
Self::from_sorted_rows(symbol, sorted)
}
let dates = sorted.iter().map(|row| row.date).collect::<Vec<_>>(); fn from_sorted_rows(symbol: String, rows: Vec<&DailyMarketSnapshot>) -> Self {
let timestamps = sorted debug_assert!(
.iter() rows.windows(2)
.map(|row| row.timestamp.clone()) .all(|window| window[0].date <= window[1].date)
.collect::<Vec<_>>(); );
let day_opens = sorted.iter().map(|row| row.day_open).collect::<Vec<_>>(); let row_count = rows.len();
let opens = sorted.iter().map(|row| row.open).collect::<Vec<_>>(); let mut dates = Vec::with_capacity(row_count);
let highs = sorted.iter().map(|row| row.high).collect::<Vec<_>>(); let mut timestamps = Vec::with_capacity(row_count);
let lows = sorted.iter().map(|row| row.low).collect::<Vec<_>>(); let mut day_opens = Vec::with_capacity(row_count);
let closes = sorted.iter().map(|row| row.close).collect::<Vec<_>>(); let mut opens = Vec::with_capacity(row_count);
let prev_closes = sorted.iter().map(|row| row.prev_close).collect::<Vec<_>>(); let mut highs = Vec::with_capacity(row_count);
let last_prices = sorted.iter().map(|row| row.last_price).collect::<Vec<_>>(); let mut lows = Vec::with_capacity(row_count);
let bid1s = sorted.iter().map(|row| row.bid1).collect::<Vec<_>>(); let mut closes = Vec::with_capacity(row_count);
let ask1s = sorted.iter().map(|row| row.ask1).collect::<Vec<_>>(); let mut prev_closes = Vec::with_capacity(row_count);
let volumes = sorted.iter().map(|row| row.volume).collect::<Vec<_>>(); let mut last_prices = Vec::with_capacity(row_count);
let minute_volumes = sorted let mut bid1s = Vec::with_capacity(row_count);
.iter() let mut ask1s = Vec::with_capacity(row_count);
.map(|row| row.minute_volume) let mut volumes = Vec::with_capacity(row_count);
.collect::<Vec<_>>(); let mut minute_volumes = Vec::with_capacity(row_count);
let bid1_volumes = sorted.iter().map(|row| row.bid1_volume).collect::<Vec<_>>(); let mut bid1_volumes = Vec::with_capacity(row_count);
let ask1_volumes = sorted.iter().map(|row| row.ask1_volume).collect::<Vec<_>>(); let mut ask1_volumes = Vec::with_capacity(row_count);
let trading_phases = sorted let mut trading_phases = Vec::with_capacity(row_count);
.iter() let mut paused = Vec::with_capacity(row_count);
.map(|row| row.trading_phase.clone()) let mut upper_limits = Vec::with_capacity(row_count);
.collect::<Vec<_>>(); let mut lower_limits = Vec::with_capacity(row_count);
let paused = sorted.iter().map(|row| row.paused).collect::<Vec<_>>(); let mut price_ticks = Vec::with_capacity(row_count);
let upper_limits = sorted.iter().map(|row| row.upper_limit).collect::<Vec<_>>(); for row in rows {
let lower_limits = sorted.iter().map(|row| row.lower_limit).collect::<Vec<_>>(); dates.push(row.date);
let price_ticks = sorted.iter().map(|row| row.price_tick).collect::<Vec<_>>(); timestamps.push(row.timestamp.clone());
day_opens.push(row.day_open);
opens.push(row.open);
highs.push(row.high);
lows.push(row.low);
closes.push(row.close);
prev_closes.push(row.prev_close);
last_prices.push(row.last_price);
bid1s.push(row.bid1);
ask1s.push(row.ask1);
volumes.push(row.volume);
minute_volumes.push(row.minute_volume);
bid1_volumes.push(row.bid1_volume);
ask1_volumes.push(row.ask1_volume);
trading_phases.push(row.trading_phase.clone());
paused.push(row.paused);
upper_limits.push(row.upper_limit);
lower_limits.push(row.lower_limit);
price_ticks.push(row.price_tick);
}
let open_prefix = prefix_sums(&opens); let open_prefix = prefix_sums(&opens);
let close_prefix = prefix_sums(&closes); let close_prefix = prefix_sums(&closes);
let prev_close_prefix = prefix_sums(&prev_closes); let prev_close_prefix = prefix_sums(&prev_closes);
@@ -1316,16 +1354,27 @@ impl DataSet {
let market_series_by_symbol = market_rows_by_symbol let market_series_by_symbol = market_rows_by_symbol
.into_par_iter() .into_par_iter()
.map(|(symbol, rows)| { .map(|(symbol, rows)| {
let series = Arc::new(SymbolPriceSeries::new(symbol.clone(), rows)); let series = Arc::new(SymbolPriceSeries::from_sorted_rows(symbol.clone(), rows));
(symbol, series) (symbol, series)
}) })
.collect::<Vec<_>>() .collect::<Vec<_>>()
.into_iter() .into_iter()
.collect::<AHashMap<_, _>>(); .collect::<AHashMap<_, _>>();
let mut factor_rows_by_symbol = AHashMap::<&str, Vec<&DailyFactorSnapshot>>::new();
for row in factor_by_date.values().flatten() {
factor_rows_by_symbol
.entry(row.symbol.as_str())
.or_default()
.push(row);
}
let adjusted_close_series_by_symbol = market_series_by_symbol let adjusted_close_series_by_symbol = market_series_by_symbol
.par_iter() .par_iter()
.filter_map(|(symbol, market)| { .filter_map(|(symbol, market)| {
AdjustedCloseSeries::new(market, &factor_by_date) let factor_rows = factor_rows_by_symbol
.get(symbol.as_str())
.map(Vec::as_slice)
.unwrap_or_default();
AdjustedCloseSeries::new(market, factor_rows)
.map(|series| (symbol.clone(), Arc::new(series))) .map(|series| (symbol.clone(), Arc::new(series)))
}) })
.collect::<Vec<_>>() .collect::<Vec<_>>()
@@ -3132,8 +3181,8 @@ fn industry_name_factor_aliases(source: &str, level: usize) -> Vec<String> {
} }
fn factor_numeric_value(snapshot: &DailyFactorSnapshot, field: &str) -> Option<f64> { fn factor_numeric_value(snapshot: &DailyFactorSnapshot, field: &str) -> Option<f64> {
let field = normalize_field(field); let field = normalized_field(field);
match field.as_str() { match field.as_ref() {
"market_cap" | "market_cap_bn" => Some(snapshot.market_cap_bn), "market_cap" | "market_cap_bn" => Some(snapshot.market_cap_bn),
"free_float_cap" | "free_float_market_cap" | "free_float_cap_bn" => { "free_float_cap" | "free_float_market_cap" | "free_float_cap_bn" => {
Some(snapshot.free_float_cap_bn) Some(snapshot.free_float_cap_bn)
@@ -3148,22 +3197,22 @@ fn factor_numeric_value(snapshot: &DailyFactorSnapshot, field: &str) -> Option<f
"effective_turnover_ratio" => snapshot.effective_turnover_ratio, "effective_turnover_ratio" => snapshot.effective_turnover_ratio,
"ths_market_value_stock" | "ths_market_value_stock_bn" => snapshot "ths_market_value_stock" | "ths_market_value_stock_bn" => snapshot
.extra_factors .extra_factors
.get(field.as_str()) .get(field.as_ref())
.copied() .copied()
.or(Some(snapshot.market_cap_bn)), .or(Some(snapshot.market_cap_bn)),
"ths_current_mv_stock" | "ths_current_mv_stock_bn" => snapshot "ths_current_mv_stock" | "ths_current_mv_stock_bn" => snapshot
.extra_factors .extra_factors
.get(field.as_str()) .get(field.as_ref())
.copied() .copied()
.or(Some(snapshot.free_float_cap_bn)), .or(Some(snapshot.free_float_cap_bn)),
"ths_turnover_ratio_stock" => snapshot "ths_turnover_ratio_stock" => snapshot
.extra_factors .extra_factors
.get(field.as_str()) .get(field.as_ref())
.copied() .copied()
.or(snapshot.turnover_ratio), .or(snapshot.turnover_ratio),
"ths_vaild_turnover_stock" | "ths_valid_turnover_stock" => snapshot "ths_vaild_turnover_stock" | "ths_valid_turnover_stock" => snapshot
.extra_factors .extra_factors
.get(field.as_str()) .get(field.as_ref())
.copied() .copied()
.or(snapshot.effective_turnover_ratio), .or(snapshot.effective_turnover_ratio),
other => snapshot.extra_factors.get(other).copied(), other => snapshot.extra_factors.get(other).copied(),
@@ -3171,7 +3220,7 @@ fn factor_numeric_value(snapshot: &DailyFactorSnapshot, field: &str) -> Option<f
} }
fn intraday_quote_numeric_value(snapshot: &IntradayExecutionQuote, field: &str) -> Option<f64> { fn intraday_quote_numeric_value(snapshot: &IntradayExecutionQuote, field: &str) -> Option<f64> {
match normalize_field(field).as_str() { match normalized_field(field).as_ref() {
"last" | "last_price" | "close" | "price" => Some(snapshot.last_price), "last" | "last_price" | "close" | "price" => Some(snapshot.last_price),
"bid1" => Some(snapshot.bid1), "bid1" => Some(snapshot.bid1),
"ask1" => Some(snapshot.ask1), "ask1" => Some(snapshot.ask1),
@@ -3904,6 +3953,38 @@ mod tests {
)); ));
} }
#[test]
fn factor_numeric_value_normalizes_fields_without_changing_aliases() {
let snapshot = DailyFactorSnapshot {
date: NaiveDate::parse_from_str("2025-01-02", "%Y-%m-%d").unwrap(),
symbol: "000001.SZ".to_string(),
market_cap_bn: 12.5,
free_float_cap_bn: 8.0,
pe_ttm: 10.0,
turnover_ratio: None,
effective_turnover_ratio: None,
extra_factors: BTreeMap::from([("custom_factor".into(), 3.5)]),
};
assert_eq!(factor_numeric_value(&snapshot, " MARKET_CAP "), Some(12.5));
assert_eq!(factor_numeric_value(&snapshot, "CUSTOM_FACTOR"), Some(3.5));
}
#[test]
fn symbol_price_series_test_constructor_sorts_unsorted_rows() {
let series = SymbolPriceSeries::new(
"000001.SZ".to_string(),
&[
market_row("2025-01-06", 12.0, 300),
market_row("2025-01-02", 10.0, 100),
market_row("2025-01-03", 11.0, 200),
],
);
assert!(series.dates.windows(2).all(|window| window[0] < window[1]));
assert_eq!(series.closes, vec![10.0, 11.0, 12.0]);
}
#[test] #[test]
fn decision_volume_average_uses_previous_completed_days_only() { fn decision_volume_average_uses_previous_completed_days_only() {
let series = SymbolPriceSeries::new( let series = SymbolPriceSeries::new(