antecedent_core/
temporal.rs1use core::fmt;
8
9use crate::ids::{Lag, VariableId};
10
11#[derive(Clone, Debug, Eq, PartialEq)]
13pub enum TemporalIndexError {
14 Invalid {
16 message: &'static str,
18 },
19 UnknownVariable {
21 id: VariableId,
23 },
24}
25
26impl fmt::Display for TemporalIndexError {
27 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
28 match self {
29 Self::Invalid { message } => write!(f, "{message}"),
30 Self::UnknownVariable { id } => write!(f, "unknown variable id {}", id.raw()),
31 }
32 }
33}
34
35impl std::error::Error for TemporalIndexError {}
36
37#[derive(Clone, Copy, Debug, Eq, PartialEq, Ord, PartialOrd, Hash)]
39pub struct TemporalNodeKey {
40 pub variable: VariableId,
42 pub offset: i32,
44}
45
46impl TemporalNodeKey {
47 #[must_use]
49 pub const fn contemporaneous(variable: VariableId) -> Self {
50 Self { variable, offset: 0 }
51 }
52
53 #[must_use]
55 pub fn lagged(variable: VariableId, lag: Lag) -> Option<Self> {
56 let lag_i = i32::try_from(lag.raw()).ok()?;
57 Some(Self { variable, offset: -lag_i })
58 }
59}
60
61#[derive(Clone, Debug, Eq, PartialEq)]
66pub struct TemporalIndexer {
67 variable_count: u32,
68 history: u32,
70 horizon: u32,
72 dense_len: usize,
74}
75
76impl TemporalIndexer {
77 pub fn new(
83 variable_count: u32,
84 history: u32,
85 horizon: u32,
86 ) -> Result<Self, TemporalIndexError> {
87 if variable_count == 0 || horizon == 0 {
88 return Err(TemporalIndexError::Invalid {
89 message: "variable_count and horizon must be non-zero",
90 });
91 }
92 let slices = history
93 .checked_add(horizon)
94 .ok_or(TemporalIndexError::Invalid { message: "history+horizon overflow" })?;
95 let dense_u32 = variable_count
96 .checked_mul(slices)
97 .ok_or(TemporalIndexError::Invalid { message: "dense index space overflow" })?;
98 let dense_len = usize::try_from(dense_u32)
99 .map_err(|_| TemporalIndexError::Invalid { message: "dense index space overflow" })?;
100 Ok(Self { variable_count, history, horizon, dense_len })
101 }
102
103 #[must_use]
105 pub const fn variable_count(&self) -> u32 {
106 self.variable_count
107 }
108
109 #[must_use]
111 pub const fn history(&self) -> u32 {
112 self.history
113 }
114
115 #[must_use]
117 pub const fn horizon(&self) -> u32 {
118 self.horizon
119 }
120
121 #[must_use]
123 pub const fn dense_len(&self) -> usize {
124 self.dense_len
125 }
126
127 pub fn dense_id(&self, key: TemporalNodeKey) -> Result<u32, TemporalIndexError> {
133 let v = key.variable.raw();
134 if v >= self.variable_count {
135 return Err(TemporalIndexError::UnknownVariable { id: key.variable });
136 }
137 let slice = i64::from(key.offset) + i64::from(self.history);
138 if slice < 0 || slice >= i64::from(self.history) + i64::from(self.horizon) {
139 return Err(TemporalIndexError::Invalid {
140 message: "temporal offset outside unfolding window",
141 });
142 }
143 let slice_u = u64::try_from(slice).map_err(|_| TemporalIndexError::Invalid {
144 message: "temporal offset outside unfolding window",
145 })?;
146 let dense = slice_u * u64::from(self.variable_count) + u64::from(v);
147 u32::try_from(dense)
148 .map_err(|_| TemporalIndexError::Invalid { message: "dense id exceeds u32" })
149 }
150
151 pub fn key_of(&self, dense: u32) -> Result<TemporalNodeKey, TemporalIndexError> {
157 let dense_usize = usize::try_from(dense)
158 .map_err(|_| TemporalIndexError::Invalid { message: "dense id out of range" })?;
159 if dense_usize >= self.dense_len() {
160 return Err(TemporalIndexError::Invalid { message: "dense id out of range" });
161 }
162 let vc = self.variable_count;
163 let slice = dense / vc;
164 let var = dense % vc;
165 let offset = i32::try_from(i64::from(slice) - i64::from(self.history))
166 .map_err(|_| TemporalIndexError::Invalid { message: "offset overflow" })?;
167 Ok(TemporalNodeKey { variable: VariableId::from_raw(var), offset })
168 }
169}
170
171#[cfg(test)]
172mod tests {
173 use super::*;
174
175 #[test]
176 fn time_major_dense_round_trip() {
177 let idx = TemporalIndexer::new(3, 2, 4).unwrap();
178 assert_eq!(idx.dense_len(), 18);
179 let key = TemporalNodeKey { variable: VariableId::from_raw(1), offset: -1 };
180 let dense = idx.dense_id(key).unwrap();
181 assert_eq!(dense, 4);
182 assert_eq!(idx.key_of(dense).unwrap(), key);
183 }
184}