antecedent_data/
multi_env_plan.rs1use std::collections::HashMap;
6use std::sync::Arc;
7
8#[cfg(test)]
9use crate::dataset::TimeSeriesData;
10use crate::error::DataError;
11use crate::multi_env::MultiEnvironmentData;
12use crate::panel::PanelData;
13use crate::reference::ReferencePointPolicy;
14use crate::sample::{LagMap, LaggedColumn, LaggedSamplePlan};
15use crate::table::TableView;
16
17pub fn plans_for_series_lengths(
23 lengths: impl IntoIterator<Item = usize>,
24 max_lag: u32,
25 columns: &Arc<[LaggedColumn]>,
26) -> Result<Vec<LaggedSamplePlan>, DataError> {
27 if columns.is_empty() {
28 return Err(DataError::InvalidArgument { message: "sample plan needs ≥1 column".into() });
29 }
30 let mut by_len: HashMap<usize, Arc<LagMap>> = HashMap::new();
31 let mut plans = Vec::new();
32 for series_len in lengths {
33 let lag_map = if let Some(m) = by_len.get(&series_len) {
34 Arc::clone(m)
35 } else {
36 let m = Arc::new(LagMap::with_reference(
37 series_len,
38 max_lag,
39 ReferencePointPolicy::SeriesOrigin,
40 )?);
41 by_len.insert(series_len, Arc::clone(&m));
42 m
43 };
44 plans.push(LaggedSamplePlan::with_shared(lag_map, Arc::clone(columns))?);
45 }
46 Ok(plans)
47}
48
49#[derive(Clone, Debug)]
54pub struct MultiEnvSamplePlan {
55 pub columns: Arc<[LaggedColumn]>,
57 pub plans: Arc<[LaggedSamplePlan]>,
59}
60
61impl MultiEnvSamplePlan {
62 pub fn try_from_multi_env(
68 data: &MultiEnvironmentData,
69 max_lag: u32,
70 columns: impl Into<Arc<[LaggedColumn]>>,
71 ) -> Result<Self, DataError> {
72 let columns = columns.into();
73 let n_env = data.env_count();
74 if n_env == 0 {
75 return Err(DataError::InvalidArgument {
76 message: "multi-env sample plan needs ≥1 environment".into(),
77 });
78 }
79 let lengths = (0..n_env).map(|i| data.environment(i).map(TableView::row_count));
80 let lengths: Result<Vec<_>, _> = lengths.collect();
81 let plans = plans_for_series_lengths(lengths?, max_lag, &columns)?;
82 Ok(Self { columns, plans: Arc::from(plans) })
83 }
84
85 #[must_use]
87 pub fn env_count(&self) -> usize {
88 self.plans.len()
89 }
90
91 pub fn plan(&self, i: usize) -> Result<&LaggedSamplePlan, DataError> {
97 self.plans.get(i).ok_or(DataError::InvalidArgument {
98 message: "multi-env plan index out of range".into(),
99 })
100 }
101}
102
103#[derive(Clone, Debug)]
105pub struct PanelSamplePlan {
106 pub columns: Arc<[LaggedColumn]>,
108 pub plans: Arc<[LaggedSamplePlan]>,
110}
111
112impl PanelSamplePlan {
113 pub fn try_from_panel(
119 panel: &PanelData,
120 max_lag: u32,
121 columns: impl Into<Arc<[LaggedColumn]>>,
122 ) -> Result<Self, DataError> {
123 let columns = columns.into();
124 if panel.unit_count() == 0 {
125 return Err(DataError::InvalidArgument {
126 message: "panel sample plan needs ≥1 unit".into(),
127 });
128 }
129 let lengths = (0..panel.unit_count()).map(|i| panel.unit(i).map(|u| u.series.row_count()));
130 let lengths: Result<Vec<_>, _> = lengths.collect();
131 let plans = plans_for_series_lengths(lengths?, max_lag, &columns)?;
132 Ok(Self { columns, plans: Arc::from(plans) })
133 }
134
135 #[must_use]
137 pub fn unit_count(&self) -> usize {
138 self.plans.len()
139 }
140}
141
142#[cfg(test)]
144#[must_use]
145pub(crate) fn series_columnar_ptr(series: &TimeSeriesData) -> *const [crate::column::OwnedColumn] {
146 series.columnar_ptr()
147}
148
149#[cfg(test)]
150mod tests {
151 use std::sync::Arc;
152
153 use antecedent_core::{Lag, VariableId};
154
155 use super::*;
156 use crate::panel::PanelUnit;
157 use crate::testing::float_series;
158
159 fn one_col() -> Arc<[LaggedColumn]> {
160 Arc::from([LaggedColumn { variable: VariableId::from_raw(0), lag: Lag::CONTEMPORANEOUS }])
161 }
162
163 #[test]
164 fn multi_env_plan_shares_geometry_without_cloning_series() {
165 let a = float_series(40, 2);
166 let b = float_series(50, 2);
167 let c = float_series(40, 2); let ptr_a = series_columnar_ptr(&a);
169 let ptr_b = series_columnar_ptr(&b);
170 let ptr_c = series_columnar_ptr(&c);
171 let multi = MultiEnvironmentData::try_new(Arc::from([a, b, c])).unwrap();
172
173 let plan = MultiEnvSamplePlan::try_from_multi_env(&multi, 2, one_col()).unwrap();
174 assert_eq!(plan.env_count(), 3);
175 assert!(Arc::ptr_eq(plan.plans[0].columns_arc(), plan.plans[1].columns_arc()));
176 assert!(Arc::ptr_eq(plan.plans[0].lag_map_arc(), plan.plans[2].lag_map_arc()));
177 assert!(!Arc::ptr_eq(plan.plans[0].lag_map_arc(), plan.plans[1].lag_map_arc()));
178
179 assert_eq!(series_columnar_ptr(multi.environment(0).unwrap()), ptr_a);
181 assert_eq!(series_columnar_ptr(multi.environment(1).unwrap()), ptr_b);
182 assert_eq!(series_columnar_ptr(multi.environment(2).unwrap()), ptr_c);
183 }
184
185 #[test]
186 fn panel_plan_shares_columns() {
187 let panel = PanelData::try_new(Arc::from([
188 PanelUnit { unit_id: 0, series: float_series(30, 2) },
189 PanelUnit { unit_id: 1, series: float_series(30, 2) },
190 ]))
191 .unwrap();
192 let plan = PanelSamplePlan::try_from_panel(&panel, 1, one_col()).unwrap();
193 assert_eq!(plan.unit_count(), 2);
194 assert!(Arc::ptr_eq(plan.plans[0].columns_arc(), plan.plans[1].columns_arc()));
195 assert!(Arc::ptr_eq(plan.plans[0].lag_map_arc(), plan.plans[1].lag_map_arc()));
196 }
197}