1#![allow(clippy::cast_possible_truncation, clippy::cast_precision_loss)]
9
10use std::sync::Arc;
11
12use antecedent_core::{KernelPolicy, VariableId};
13
14use crate::error::DataError;
15use crate::lagged_frame::LaggedFrame;
16use crate::multi_env::MultiEnvironmentData;
17use crate::table::TableView;
18
19pub const DEFAULT_MAX_TIME_ONE_HOT_LEVELS: usize = 512;
21
22#[derive(Clone, Copy, Debug, Default, Eq, PartialEq, Hash)]
24pub enum TimeDummyEncoding {
25 #[default]
27 IntegerIndex,
28 OneHot,
32}
33
34#[derive(Clone, Copy, Debug, Eq, PartialEq)]
36pub struct DummyOptions {
37 pub include_space_dummy: bool,
39 pub include_time_dummy: bool,
41 pub time_dummy_encoding: TimeDummyEncoding,
43 pub max_time_one_hot_levels: usize,
45}
46
47impl Default for DummyOptions {
48 fn default() -> Self {
49 Self {
50 include_space_dummy: true,
51 include_time_dummy: false,
52 time_dummy_encoding: TimeDummyEncoding::IntegerIndex,
53 max_time_one_hot_levels: DEFAULT_MAX_TIME_ONE_HOT_LEVELS,
54 }
55 }
56}
57
58#[derive(Clone, Debug)]
60pub struct PooledLaggedFrame {
61 pub frame: LaggedFrame,
63 pub observed_variables: Arc<[VariableId]>,
65 pub space_dummy_variables: Arc<[VariableId]>,
67 pub time_dummy_variables: Arc<[VariableId]>,
69 pub env_effective_rows: Arc<[usize]>,
71}
72
73impl PooledLaggedFrame {
74 #[must_use]
76 pub fn all_variables(&self) -> Vec<VariableId> {
77 let mut v = self.observed_variables.to_vec();
78 v.extend_from_slice(&self.space_dummy_variables);
79 v.extend_from_slice(&self.time_dummy_variables);
80 v
81 }
82
83 #[must_use]
85 pub fn is_dummy(&self, v: VariableId) -> bool {
86 self.space_dummy_variables.iter().any(|&x| x == v)
87 || self.time_dummy_variables.iter().any(|&x| x == v)
88 }
89}
90
91pub fn pool_multi_env_lagged_frame(
102 data: &MultiEnvironmentData,
103 variables: &[VariableId],
104 frame_depth: u32,
105 dummies: DummyOptions,
106 policy: &KernelPolicy,
107) -> Result<PooledLaggedFrame, DataError> {
108 if data.env_count() == 0 {
109 return Err(DataError::InvalidArgument {
110 message: "pooled lagged frame needs ≥1 environment".into(),
111 });
112 }
113 if variables.is_empty() {
114 return Err(DataError::InvalidArgument {
115 message: "pooled lagged frame needs ≥1 variable".into(),
116 });
117 }
118
119 let mut frames = Vec::with_capacity(data.env_count());
120 let mut env_effective = Vec::with_capacity(data.env_count());
121 for i in 0..data.env_count() {
122 let series = data.environment(i)?;
123 let frame = LaggedFrame::from_series(series, variables, frame_depth, policy)?;
124 env_effective.push(frame.n_effective());
125 frames.push(frame);
126 }
127 let mut pooled = LaggedFrame::stack(&frames)?;
128
129 let next_id = next_synthetic_id(variables);
130 let mut space_dummies = Vec::new();
131 let mut time_dummies = Vec::new();
132 let mut cursor = next_id;
133
134 if dummies.include_space_dummy && data.env_count() > 1 {
135 let m = data.env_count();
136 let n_hot = m - 1;
137 let mut cols: Vec<(VariableId, Vec<f64>)> = Vec::with_capacity(n_hot);
138 for k in 0..n_hot {
139 let id = VariableId::from_raw(cursor);
140 cursor = cursor.saturating_add(1);
141 space_dummies.push(id);
142 let mut col = Vec::with_capacity(pooled.n_effective());
143 for (env_i, &n_eff) in env_effective.iter().enumerate() {
144 let val = if env_i == k { 1.0 } else { 0.0 };
145 col.extend(std::iter::repeat_n(val, n_eff));
146 }
147 cols.push((id, col));
148 }
149 pooled = pooled.append_constant_lag_columns(&cols)?;
150 }
151
152 if dummies.include_time_dummy {
153 let times = effective_raw_times(data, &env_effective, frame_depth)?;
154 match dummies.time_dummy_encoding {
155 TimeDummyEncoding::IntegerIndex => {
156 let id = VariableId::from_raw(cursor);
157 cursor = cursor.saturating_add(1);
158 time_dummies.push(id);
159 let col: Vec<f64> = times.iter().map(|&t| f64::from(t)).collect();
160 pooled = pooled.append_constant_lag_columns(&[(id, col)])?;
161 }
162 TimeDummyEncoding::OneHot => {
163 let mut levels = times.clone();
164 levels.sort_unstable();
165 levels.dedup();
166 if levels.len() > dummies.max_time_one_hot_levels {
167 return Err(DataError::InvalidArgument {
168 message: format!(
169 "time one-hot: {} distinct levels exceeds max_time_one_hot_levels={}",
170 levels.len(),
171 dummies.max_time_one_hot_levels
172 ),
173 });
174 }
175 if levels.len() > 1 {
177 let n_hot = levels.len() - 1;
178 let mut cols: Vec<(VariableId, Vec<f64>)> = Vec::with_capacity(n_hot);
179 for &level in &levels[..n_hot] {
180 let id = VariableId::from_raw(cursor);
181 cursor = cursor.saturating_add(1);
182 time_dummies.push(id);
183 let col: Vec<f64> =
184 times.iter().map(|&t| if t == level { 1.0 } else { 0.0 }).collect();
185 cols.push((id, col));
186 }
187 pooled = pooled.append_constant_lag_columns(&cols)?;
188 }
189 }
190 }
191 }
192
193 let _ = cursor; Ok(PooledLaggedFrame {
195 frame: pooled,
196 observed_variables: Arc::from(variables.to_vec()),
197 space_dummy_variables: Arc::from(space_dummies),
198 time_dummy_variables: Arc::from(time_dummies),
199 env_effective_rows: Arc::from(env_effective),
200 })
201}
202
203fn effective_raw_times(
205 data: &MultiEnvironmentData,
206 env_effective: &[usize],
207 frame_depth: u32,
208) -> Result<Vec<u32>, DataError> {
209 let mut times = Vec::new();
210 let base = frame_depth as usize;
211 for (env_idx, &n_eff) in env_effective.iter().enumerate() {
212 let series = data.environment(env_idx)?;
213 let n_raw = series.row_count();
214 if base + n_eff > n_raw {
215 return Err(DataError::InvalidArgument {
216 message: format!(
217 "time dummy: env {env_idx} effective {n_eff} + depth {frame_depth} exceeds rows {n_raw}"
218 ),
219 });
220 }
221 for j in 0..n_eff {
222 times.push(u32::try_from(base + j).map_err(|_| DataError::InvalidArgument {
223 message: "time dummy: raw time index exceeds u32".into(),
224 })?);
225 }
226 }
227 Ok(times)
228}
229
230fn next_synthetic_id(variables: &[VariableId]) -> u32 {
231 variables.iter().map(|v| v.raw()).max().map_or(0, |m| m.saturating_add(1))
232}
233
234#[cfg(test)]
235#[allow(clippy::cast_precision_loss)]
236mod tests {
237 use std::sync::Arc;
238
239 use antecedent_core::{Lag, VariableId};
240
241 use super::*;
242 use crate::multi_env::MultiEnvironmentData;
243 use crate::testing::float_series;
244
245 #[test]
246 fn pool_stacks_effective_rows_without_cross_env_bleed() {
247 let a = float_series(20, 2);
248 let b = float_series(30, 2);
249 let multi = MultiEnvironmentData::try_new(Arc::from([a, b])).unwrap();
250 let vars = [VariableId::from_raw(0), VariableId::from_raw(1)];
251 let depth = 2u32;
252 let pooled = pool_multi_env_lagged_frame(
253 &multi,
254 &vars,
255 depth,
256 DummyOptions {
257 include_space_dummy: false,
258 include_time_dummy: false,
259 ..DummyOptions::default()
260 },
261 &KernelPolicy::default_policy(),
262 )
263 .unwrap();
264 assert_eq!(pooled.frame.n_effective(), 46);
266 assert_eq!(pooled.env_effective_rows.as_ref(), &[18, 28]);
267 assert!(pooled.space_dummy_variables.is_empty());
268 assert!(pooled.time_dummy_variables.is_empty());
269 }
270
271 #[test]
272 fn space_dummy_one_hot_m_minus_1() {
273 let a = float_series(16, 2);
274 let b = float_series(16, 2);
275 let c = float_series(16, 2);
276 let multi = MultiEnvironmentData::try_new(Arc::from([a, b, c])).unwrap();
277 let vars = [VariableId::from_raw(0), VariableId::from_raw(1)];
278 let pooled = pool_multi_env_lagged_frame(
279 &multi,
280 &vars,
281 2,
282 DummyOptions {
283 include_space_dummy: true,
284 include_time_dummy: false,
285 ..DummyOptions::default()
286 },
287 &KernelPolicy::default_policy(),
288 )
289 .unwrap();
290 assert_eq!(pooled.space_dummy_variables.len(), 2);
291 let d0 = pooled.space_dummy_variables[0];
292 let col = pooled.frame.column(pooled.frame.column_index(d0, Lag::CONTEMPORANEOUS).unwrap());
293 assert!((col[0] - 1.0).abs() < 1e-12);
295 assert!((col[13] - 1.0).abs() < 1e-12);
296 assert!((col[14]).abs() < 1e-12);
298 }
299
300 #[test]
301 fn time_dummy_integer_tracks_raw_time() {
302 let a = float_series(12, 1);
303 let multi = MultiEnvironmentData::try_new(Arc::from([a])).unwrap();
304 let vars = [VariableId::from_raw(0)];
305 let pooled = pool_multi_env_lagged_frame(
306 &multi,
307 &vars,
308 2,
309 DummyOptions {
310 include_space_dummy: false,
311 include_time_dummy: true,
312 time_dummy_encoding: TimeDummyEncoding::IntegerIndex,
313 ..DummyOptions::default()
314 },
315 &KernelPolicy::default_policy(),
316 )
317 .unwrap();
318 assert_eq!(pooled.time_dummy_variables.len(), 1);
319 let tid = pooled.time_dummy_variables[0];
320 let col =
321 pooled.frame.column(pooled.frame.column_index(tid, Lag::CONTEMPORANEOUS).unwrap());
322 assert_eq!(col.len(), 10);
323 assert!((col[0] - 2.0).abs() < 1e-12);
324 assert!((col[9] - 11.0).abs() < 1e-12);
325 }
326
327 #[test]
328 fn time_dummy_one_hot_t_minus_1() {
329 let a = float_series(6, 1);
331 let multi = MultiEnvironmentData::try_new(Arc::from([a])).unwrap();
332 let vars = [VariableId::from_raw(0)];
333 let pooled = pool_multi_env_lagged_frame(
334 &multi,
335 &vars,
336 2,
337 DummyOptions {
338 include_space_dummy: false,
339 include_time_dummy: true,
340 time_dummy_encoding: TimeDummyEncoding::OneHot,
341 ..DummyOptions::default()
342 },
343 &KernelPolicy::default_policy(),
344 )
345 .unwrap();
346 assert_eq!(pooled.time_dummy_variables.len(), 3);
347 let t0 = pooled.time_dummy_variables[0];
348 let col0 =
349 pooled.frame.column(pooled.frame.column_index(t0, Lag::CONTEMPORANEOUS).unwrap());
350 assert!((col0[0] - 1.0).abs() < 1e-12);
352 assert!(col0[1].abs() < 1e-12);
354 let last = col0.len() - 1;
356 for &tid in pooled.time_dummy_variables.iter() {
357 let col =
358 pooled.frame.column(pooled.frame.column_index(tid, Lag::CONTEMPORANEOUS).unwrap());
359 assert!(col[last].abs() < 1e-12, "reference time should be all-zero");
360 }
361 }
362
363 #[test]
364 fn time_dummy_one_hot_respects_level_cap() {
365 let a = float_series(20, 1);
366 let multi = MultiEnvironmentData::try_new(Arc::from([a])).unwrap();
367 let vars = [VariableId::from_raw(0)];
368 let err = pool_multi_env_lagged_frame(
369 &multi,
370 &vars,
371 2,
372 DummyOptions {
373 include_space_dummy: false,
374 include_time_dummy: true,
375 time_dummy_encoding: TimeDummyEncoding::OneHot,
376 max_time_one_hot_levels: 4,
377 },
378 &KernelPolicy::default_policy(),
379 )
380 .unwrap_err();
381 assert!(err.to_string().contains("max_time_one_hot_levels"), "unexpected error: {err}");
382 }
383
384 #[test]
385 fn time_dummy_one_hot_aligns_across_unequal_envs() {
386 let a = float_series(6, 1);
389 let b = float_series(8, 1);
390 let multi = MultiEnvironmentData::try_new(Arc::from([a, b])).unwrap();
391 let vars = [VariableId::from_raw(0)];
392 let pooled = pool_multi_env_lagged_frame(
393 &multi,
394 &vars,
395 2,
396 DummyOptions {
397 include_space_dummy: false,
398 include_time_dummy: true,
399 time_dummy_encoding: TimeDummyEncoding::OneHot,
400 ..DummyOptions::default()
401 },
402 &KernelPolicy::default_policy(),
403 )
404 .unwrap();
405 assert_eq!(pooled.frame.n_effective(), 4 + 6);
406 assert_eq!(pooled.time_dummy_variables.len(), 5);
407 }
408}