Skip to main content

antecedent_data/
pooled_frame.rs

1//! Pooled multi-environment lagged frames with space/time dummies (J-PCMCI+).
2//!
3//! Materializes each environment independently (no cross-env lag windows), stacks
4//! effective rows, then appends synthetic dummy columns.
5//!
6//! SPDX-License-Identifier: MIT OR Apache-2.0
7
8#![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
19/// Default cap on distinct time levels for [`TimeDummyEncoding::OneHot`] (fail-closed).
20pub const DEFAULT_MAX_TIME_ONE_HOT_LEVELS: usize = 512;
21
22/// How the synthetic time dummy is embedded in the pooled frame.
23#[derive(Clone, Copy, Debug, Default, Eq, PartialEq, Hash)]
24pub enum TimeDummyEncoding {
25    /// Single integer raw-time index column (legacy opt-in).
26    #[default]
27    IntegerIndex,
28    /// One-hot over distinct raw times in the pooled effective rows (`T−1` columns;
29    /// last level is the reference category — same convention as space dummies).
30    /// Paper-faithful Günther embedding (up to the dropped reference column).
31    OneHot,
32}
33
34/// Options for synthesizing dataset / time dummy columns on a pooled frame.
35#[derive(Clone, Copy, Debug, Eq, PartialEq)]
36pub struct DummyOptions {
37    /// One-hot space (dataset) dummy with `M−1` columns when `M > 1`.
38    pub include_space_dummy: bool,
39    /// When true, synthesize a time dummy under [`Self::time_dummy_encoding`].
40    pub include_time_dummy: bool,
41    /// Integer index vs one-hot of `T` (ignored unless [`Self::include_time_dummy`]).
42    pub time_dummy_encoding: TimeDummyEncoding,
43    /// Fail-closed cap on distinct time levels for [`TimeDummyEncoding::OneHot`].
44    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/// Result of pooling multi-environment series into one lagged frame.
59#[derive(Clone, Debug)]
60pub struct PooledLaggedFrame {
61    /// Stacked frame including system/context variables and any dummy columns.
62    pub frame: LaggedFrame,
63    /// Variable ids that were requested for pooling (system + observed context).
64    pub observed_variables: Arc<[VariableId]>,
65    /// Synthetic space-dummy variable ids (`M−1` one-hot columns), empty if off.
66    pub space_dummy_variables: Arc<[VariableId]>,
67    /// Synthetic time-dummy variable ids (one integer column, or `T−1` one-hot columns).
68    pub time_dummy_variables: Arc<[VariableId]>,
69    /// Per-environment effective row counts in stack order.
70    pub env_effective_rows: Arc<[usize]>,
71}
72
73impl PooledLaggedFrame {
74    /// All variable ids present in [`Self::frame`] (observed then space dummies then time).
75    #[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    /// Whether `v` is a synthesized dummy column.
84    #[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
91/// Build a pooled lagged frame from multi-environment data.
92///
93/// Each environment is materialized with [`LaggedFrame::from_series`] at
94/// `frame_depth`, then vertically stacked. Dummy columns are appended as
95/// contemporaneous-constant (space) or time-index / one-hot (time) series.
96///
97/// # Errors
98///
99/// Empty multi-env / variable list, per-env frame failures, stack geometry mismatch,
100/// time alignment failure, or one-hot level count above the configured cap.
101pub 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                // T≤1 → constant; no identifiable dummy columns.
176                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; // advance tracked for future synthetic columns
194    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
203/// Absolute raw time index for each stacked effective row (`j + frame_depth` under `SeriesOrigin`).
204fn 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        // n_effective = (20-2) + (30-2) = 46
265        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        // First env block: 14 effective rows of 1.0
294        assert!((col[0] - 1.0).abs() < 1e-12);
295        assert!((col[13] - 1.0).abs() < 1e-12);
296        // Second env: 0.0 for first hot column
297        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        // depth=2, n=6 → effective times {2,3,4,5} → T=4 → 3 one-hot columns.
330        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        // Row 0 → time 2 → first level → 1
351        assert!((col0[0] - 1.0).abs() < 1e-12);
352        // Row 1 → time 3 → 0 on first column
353        assert!(col0[1].abs() < 1e-12);
354        // Last row → time 5 (reference) → all zeros
355        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        // Env A: times 2..5 (n=6, depth=2); Env B: times 2..7 (n=8, depth=2).
387        // Union levels {2,3,4,5,6,7} → 5 one-hot columns.
388        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}