Skip to main content

antecedent_model/
batch.rs

1//! Columnar value / noise batches and mechanism workspaces.
2//!
3//! SPDX-License-Identifier: MIT OR Apache-2.0
4
5#![allow(clippy::cast_possible_truncation)]
6
7use std::sync::Arc;
8
9use crate::error::ModelError;
10
11/// Column-major batch of continuous values: `values[node * n_rows + row]`.
12#[derive(Clone, Debug, Default)]
13pub struct ValueBatch {
14    /// Number of rows (samples / units).
15    pub n_rows: usize,
16    /// Number of nodes (columns).
17    pub n_nodes: usize,
18    /// Flat column-major storage.
19    pub values: Arc<[f64]>,
20}
21
22impl ValueBatch {
23    /// Allocate zeros.
24    #[must_use]
25    pub fn zeros(n_rows: usize, n_nodes: usize) -> Self {
26        Self { n_rows, n_nodes, values: Arc::from(vec![0.0; n_rows.saturating_mul(n_nodes)]) }
27    }
28
29    /// Borrow a column.
30    ///
31    /// # Errors
32    ///
33    /// Out of range.
34    pub fn column(&self, node: usize) -> Result<&[f64], ModelError> {
35        if node >= self.n_nodes {
36            return Err(ModelError::Shape { message: "value column out of range".into() });
37        }
38        let start = node * self.n_rows;
39        Ok(&self.values[start..start + self.n_rows])
40    }
41
42    /// Value at `(row, node)`.
43    ///
44    /// # Errors
45    ///
46    /// Out of range.
47    pub fn get(&self, row: usize, node: usize) -> Result<f64, ModelError> {
48        if row >= self.n_rows || node >= self.n_nodes {
49            return Err(ModelError::Shape { message: "value index out of range".into() });
50        }
51        Ok(self.values[node * self.n_rows + row])
52    }
53}
54
55/// Mutable view into a value batch (owned buffer).
56#[derive(Debug)]
57pub struct ValueBatchMut<'a> {
58    /// Rows.
59    pub n_rows: usize,
60    /// Nodes.
61    pub n_nodes: usize,
62    /// Flat column-major storage.
63    pub values: &'a mut [f64],
64}
65
66impl<'a> ValueBatchMut<'a> {
67    /// Wrap a buffer.
68    ///
69    /// # Errors
70    ///
71    /// Length mismatch.
72    pub fn new(n_rows: usize, n_nodes: usize, values: &'a mut [f64]) -> Result<Self, ModelError> {
73        if values.len() < n_rows.saturating_mul(n_nodes) {
74            return Err(ModelError::Shape { message: "value buffer too short".into() });
75        }
76        Ok(Self { n_rows, n_nodes, values })
77    }
78
79    /// Mutable column slice.
80    ///
81    /// # Errors
82    ///
83    /// Out of range.
84    pub fn column_mut(&mut self, node: usize) -> Result<&mut [f64], ModelError> {
85        if node >= self.n_nodes {
86            return Err(ModelError::Shape { message: "value column out of range".into() });
87        }
88        let start = node * self.n_rows;
89        Ok(&mut self.values[start..start + self.n_rows])
90    }
91
92    /// Set `(row, node)`.
93    ///
94    /// # Errors
95    ///
96    /// Out of range.
97    pub fn set(&mut self, row: usize, node: usize, v: f64) -> Result<(), ModelError> {
98        if row >= self.n_rows || node >= self.n_nodes {
99            return Err(ModelError::Shape { message: "value index out of range".into() });
100        }
101        self.values[node * self.n_rows + row] = v;
102        Ok(())
103    }
104
105    /// Freeze into an owned [`ValueBatch`].
106    #[must_use]
107    pub fn into_batch(self) -> ValueBatch {
108        ValueBatch {
109            n_rows: self.n_rows,
110            n_nodes: self.n_nodes,
111            values: Arc::from(self.values.to_vec()),
112        }
113    }
114}
115
116/// Columnar exogenous noise batch (same layout as [`ValueBatch`]).
117#[derive(Clone, Debug, Default)]
118pub struct NoiseBatch {
119    /// Rows.
120    pub n_rows: usize,
121    /// Nodes.
122    pub n_nodes: usize,
123    /// Flat storage.
124    pub values: Arc<[f64]>,
125}
126
127impl NoiseBatch {
128    /// Zeros.
129    #[must_use]
130    pub fn zeros(n_rows: usize, n_nodes: usize) -> Self {
131        Self { n_rows, n_nodes, values: Arc::from(vec![0.0; n_rows.saturating_mul(n_nodes)]) }
132    }
133
134    /// Column.
135    ///
136    /// # Errors
137    ///
138    /// Out of range.
139    pub fn column(&self, node: usize) -> Result<&[f64], ModelError> {
140        if node >= self.n_nodes {
141            return Err(ModelError::Shape { message: "noise column out of range".into() });
142        }
143        let start = node * self.n_rows;
144        Ok(&self.values[start..start + self.n_rows])
145    }
146}
147
148/// Mutable noise batch.
149#[derive(Debug)]
150pub struct NoiseBatchMut<'a> {
151    /// Rows.
152    pub n_rows: usize,
153    /// Nodes.
154    pub n_nodes: usize,
155    /// Storage.
156    pub values: &'a mut [f64],
157}
158
159impl<'a> NoiseBatchMut<'a> {
160    /// Wrap.
161    ///
162    /// # Errors
163    ///
164    /// Length mismatch.
165    pub fn new(n_rows: usize, n_nodes: usize, values: &'a mut [f64]) -> Result<Self, ModelError> {
166        if values.len() < n_rows.saturating_mul(n_nodes) {
167            return Err(ModelError::Shape { message: "noise buffer too short".into() });
168        }
169        Ok(Self { n_rows, n_nodes, values })
170    }
171
172    /// Immutable column.
173    ///
174    /// # Errors
175    ///
176    /// Out of range.
177    pub fn column(&self, node: usize) -> Result<&[f64], ModelError> {
178        if node >= self.n_nodes {
179            return Err(ModelError::Shape { message: "noise column out of range".into() });
180        }
181        let start = node * self.n_rows;
182        Ok(&self.values[start..start + self.n_rows])
183    }
184
185    /// Mutable column.
186    ///
187    /// # Errors
188    ///
189    /// Out of range.
190    pub fn column_mut(&mut self, node: usize) -> Result<&mut [f64], ModelError> {
191        if node >= self.n_nodes {
192            return Err(ModelError::Shape { message: "noise column out of range".into() });
193        }
194        let start = node * self.n_rows;
195        Ok(&mut self.values[start..start + self.n_rows])
196    }
197
198    /// Freeze.
199    #[must_use]
200    pub fn into_batch(self) -> NoiseBatch {
201        NoiseBatch {
202            n_rows: self.n_rows,
203            n_nodes: self.n_nodes,
204            values: Arc::from(self.values.to_vec()),
205        }
206    }
207}
208
209/// Borrowed parent columns for one node (aligned row-major view into gathered parents).
210#[derive(Clone, Copy, Debug)]
211pub struct ParentBatch<'a> {
212    /// Number of rows.
213    pub n_rows: usize,
214    /// Number of parents.
215    pub n_parents: usize,
216    /// Flat `parent * n_rows + row`.
217    pub values: &'a [f64],
218}
219
220impl<'a> ParentBatch<'a> {
221    /// Empty parents.
222    #[must_use]
223    pub const fn empty(n_rows: usize) -> Self {
224        Self { n_rows, n_parents: 0, values: &[] }
225    }
226
227    /// Parent column `p`.
228    ///
229    /// # Errors
230    ///
231    /// Out of range.
232    pub fn column(&self, p: usize) -> Result<&'a [f64], ModelError> {
233        if p >= self.n_parents {
234            return Err(ModelError::Shape { message: "parent column out of range".into() });
235        }
236        let start = p * self.n_rows;
237        Ok(&self.values[start..start + self.n_rows])
238    }
239}
240
241/// Reusable scratch for mechanism evaluation.
242#[derive(Clone, Debug, Default)]
243pub struct MechanismWorkspace {
244    /// Gathered parent matrix (column-major over parents).
245    pub parents: Vec<f64>,
246    /// Scratch residuals / linear predictors.
247    pub scratch: Vec<f64>,
248    /// Grow counter (tests / reuse gates).
249    pub grow_count: u32,
250}
251
252impl MechanismWorkspace {
253    /// Ensure capacity for `n_rows` × `n_parents` gather + scratch of `n_rows`.
254    pub fn prepare(&mut self, n_rows: usize, n_parents: usize) {
255        let need = n_rows.saturating_mul(n_parents.max(1));
256        if self.parents.capacity() < need {
257            self.parents.reserve(need - self.parents.capacity());
258            self.grow_count = self.grow_count.saturating_add(1);
259        }
260        self.parents.resize(need, 0.0);
261        if self.scratch.capacity() < n_rows {
262            self.scratch.reserve(n_rows.saturating_sub(self.scratch.capacity()));
263            self.grow_count = self.grow_count.saturating_add(1);
264        }
265        self.scratch.resize(n_rows, 0.0);
266    }
267}