diff --git a/crates/fidc-core/src/data.rs b/crates/fidc-core/src/data.rs index 4fe5fcf..ccf2fa9 100644 --- a/crates/fidc-core/src/data.rs +++ b/crates/fidc-core/src/data.rs @@ -87,6 +87,15 @@ pub enum DataSetError { DuplicateIntradayMarketOverlay { date: NaiveDate, symbol: String }, #[error("cannot mutate shared {component} while finalizing a backtest dataset")] SharedComponentMutation { component: &'static str }, + #[error( + "{kind} snapshot rows and symbol ids are misaligned on {date}: rows={row_count}, ids={symbol_id_count}" + )] + SnapshotSymbolIndexAlignment { + kind: &'static str, + date: NaiveDate, + row_count: usize, + symbol_id_count: usize, + }, } #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -1780,43 +1789,88 @@ impl DataSet { .into_iter() .map(|instrument| (instrument.symbol.clone(), instrument)) .collect::>(); - let mut market_rows_by_symbol = AHashMap::>::new(); - for row in market_by_date.values().flatten() { - if let Some(rows) = market_rows_by_symbol.get_mut(row.symbol.as_str()) { - rows.push(row); - continue; + let symbol_id_by_code = build_symbol_id_index( + &instruments, + &market_by_date, + &factor_by_date, + &candidate_by_date, + ); + let symbol_count = symbol_id_by_code.len(); + let mut symbol_by_id = vec![Arc::::from(""); symbol_count]; + for (symbol, symbol_id) in &symbol_id_by_code { + symbol_by_id[*symbol_id as usize] = Arc::::from(symbol.as_str()); + } + let mut instruments_by_symbol_id = vec![None; symbol_count]; + for (symbol, instrument) in &instruments { + if let Some(symbol_id) = symbol_id_by_code.get(symbol).copied() { + instruments_by_symbol_id[symbol_id as usize] = Some(instrument.clone()); } - market_rows_by_symbol.insert(row.symbol.clone(), vec![row]); } - let market_rows_by_symbol = market_rows_by_symbol.into_iter().collect::>(); - let market_series_by_symbol = market_rows_by_symbol + let market_symbol_ids_by_date = + build_group_symbol_ids(&market_by_date, &symbol_id_by_code, |item| { + item.symbol.as_str() + }); + let factor_symbol_ids_by_date = + build_group_symbol_ids(&factor_by_date, &symbol_id_by_code, |item| { + item.symbol.as_str() + }); + let candidate_symbol_ids_by_date = + build_group_symbol_ids(&candidate_by_date, &symbol_id_by_code, |item| { + item.symbol.as_str() + }); + + let market_rows_by_symbol_id = group_rows_by_symbol_id( + "market", + &market_by_date, + &market_symbol_ids_by_date, + symbol_count, + )?; + let market_series_by_symbol_id = market_rows_by_symbol_id .into_par_iter() - .map(|(symbol, rows)| { - let series = Arc::new(SymbolPriceSeries::from_sorted_rows(symbol.clone(), rows)); - (symbol, series) + .enumerate() + .map(|(symbol_id, rows)| { + (!rows.is_empty()).then(|| { + Arc::new(SymbolPriceSeries::from_sorted_rows( + symbol_by_id[symbol_id].to_string(), + rows, + )) + }) + }) + .collect::>(); + let market_series_by_symbol = market_series_by_symbol_id + .iter() + .enumerate() + .filter_map(|(symbol_id, series)| { + series.as_ref().map(|series| { + (symbol_by_id[symbol_id].to_string(), Arc::clone(series)) + }) }) - .collect::>() - .into_iter() .collect::>(); - 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 factor_rows_by_symbol_id = group_rows_by_symbol_id( + "factor", + &factor_by_date, + &factor_symbol_ids_by_date, + symbol_count, + )?; + let adjusted_close_series_by_symbol_id = market_series_by_symbol_id .par_iter() - .filter_map(|(symbol, market)| { - 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))) + .enumerate() + .map(|(symbol_id, market)| { + market.as_ref().and_then(|market| { + AdjustedCloseSeries::new(market, &factor_rows_by_symbol_id[symbol_id]) + .map(Arc::new) + }) + }) + .collect::>(); + let adjusted_close_series_by_symbol = adjusted_close_series_by_symbol_id + .iter() + .enumerate() + .filter_map(|(symbol_id, series)| { + series.as_ref().map(|series| { + (symbol_by_id[symbol_id].to_string(), Arc::clone(series)) + }) }) - .collect::>() - .into_iter() .collect::>(); let factor_texts = factor_texts .into_iter() @@ -1835,63 +1889,23 @@ impl DataSet { .map(|item| ((item.date, item.symbol.clone(), item.field.clone()), item)) .collect::>(); - let symbol_id_by_code = build_symbol_id_index( - &instruments, - &market_by_date, - &factor_by_date, - &candidate_by_date, - ); - let mut symbol_by_id = vec![Arc::::from(""); symbol_id_by_code.len()]; - for (symbol, symbol_id) in &symbol_id_by_code { - symbol_by_id[*symbol_id as usize] = Arc::::from(symbol.as_str()); - } - let mut instruments_by_symbol_id = vec![None; symbol_id_by_code.len()]; - for (symbol, instrument) in &instruments { - if let Some(symbol_id) = symbol_id_by_code.get(symbol).copied() { - instruments_by_symbol_id[symbol_id as usize] = Some(instrument.clone()); - } - } - let market_symbol_ids_by_date = - build_group_symbol_ids(&market_by_date, &symbol_id_by_code, |item| { - item.symbol.as_str() - }); - let factor_symbol_ids_by_date = - build_group_symbol_ids(&factor_by_date, &symbol_id_by_code, |item| { - item.symbol.as_str() - }); let factor_market_cap_order_by_date = build_factor_market_cap_order(&factor_by_date, &factor_symbol_ids_by_date); - let candidate_symbol_ids_by_date = - build_group_symbol_ids(&candidate_by_date, &symbol_id_by_code, |item| { - item.symbol.as_str() - }); let market_row_positions_by_date = build_dense_row_positions( &market_by_date, &market_symbol_ids_by_date, - symbol_id_by_code.len(), + symbol_count, ); let factor_row_positions_by_date = build_dense_row_positions( &factor_by_date, &factor_symbol_ids_by_date, - symbol_id_by_code.len(), + symbol_count, ); let candidate_row_positions_by_date = build_dense_row_positions( &candidate_by_date, &candidate_symbol_ids_by_date, - symbol_id_by_code.len(), + symbol_count, ); - let mut market_series_by_symbol_id = vec![None; symbol_id_by_code.len()]; - for (symbol, series) in &market_series_by_symbol { - if let Some(symbol_id) = symbol_id_by_code.get(symbol).copied() { - market_series_by_symbol_id[symbol_id as usize] = Some(Arc::clone(series)); - } - } - let mut adjusted_close_series_by_symbol_id = vec![None; symbol_id_by_code.len()]; - for (symbol, series) in &adjusted_close_series_by_symbol { - if let Some(symbol_id) = symbol_id_by_code.get(symbol).copied() { - adjusted_close_series_by_symbol_id[symbol_id as usize] = Some(Arc::clone(series)); - } - } let market_series_end_positions_by_calendar_index = build_calendar_series_end_positions(&market_series_by_symbol_id, &calendar); let execution_quotes_by_date = build_execution_quote_index(execution_quotes); @@ -4493,6 +4507,34 @@ where .collect() } +fn group_rows_by_symbol_id<'a, T>( + kind: &'static str, + groups: &'a BTreeMap>, + symbol_ids_by_date: &BTreeMap>, + symbol_count: usize, +) -> Result>, DataSetError> { + let mut rows_by_symbol_id = (0..symbol_count) + .map(|_| Vec::<&T>::new()) + .collect::>(); + for (date, rows) in groups { + let symbol_ids = symbol_ids_by_date + .get(date) + .expect("daily snapshot symbol ids must exist before series grouping"); + if rows.len() != symbol_ids.len() { + return Err(DataSetError::SnapshotSymbolIndexAlignment { + kind, + date: *date, + row_count: rows.len(), + symbol_id_count: symbol_ids.len(), + }); + } + for (row, symbol_id) in rows.iter().zip(symbol_ids) { + rows_by_symbol_id[*symbol_id as usize].push(row); + } + } + Ok(rows_by_symbol_id) +} + fn build_factor_market_cap_order( factor_by_date: &BTreeMap>, factor_symbol_ids_by_date: &BTreeMap>,