antecedent_model/
batch.rs1#![allow(clippy::cast_possible_truncation)]
6
7use std::sync::Arc;
8
9use crate::error::ModelError;
10
11#[derive(Clone, Debug, Default)]
13pub struct ValueBatch {
14 pub n_rows: usize,
16 pub n_nodes: usize,
18 pub values: Arc<[f64]>,
20}
21
22impl ValueBatch {
23 #[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 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 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#[derive(Debug)]
57pub struct ValueBatchMut<'a> {
58 pub n_rows: usize,
60 pub n_nodes: usize,
62 pub values: &'a mut [f64],
64}
65
66impl<'a> ValueBatchMut<'a> {
67 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 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 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 #[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#[derive(Clone, Debug, Default)]
118pub struct NoiseBatch {
119 pub n_rows: usize,
121 pub n_nodes: usize,
123 pub values: Arc<[f64]>,
125}
126
127impl NoiseBatch {
128 #[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 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#[derive(Debug)]
150pub struct NoiseBatchMut<'a> {
151 pub n_rows: usize,
153 pub n_nodes: usize,
155 pub values: &'a mut [f64],
157}
158
159impl<'a> NoiseBatchMut<'a> {
160 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 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 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 #[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#[derive(Clone, Copy, Debug)]
211pub struct ParentBatch<'a> {
212 pub n_rows: usize,
214 pub n_parents: usize,
216 pub values: &'a [f64],
218}
219
220impl<'a> ParentBatch<'a> {
221 #[must_use]
223 pub const fn empty(n_rows: usize) -> Self {
224 Self { n_rows, n_parents: 0, values: &[] }
225 }
226
227 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#[derive(Clone, Debug, Default)]
243pub struct MechanismWorkspace {
244 pub parents: Vec<f64>,
246 pub scratch: Vec<f64>,
248 pub grow_count: u32,
250}
251
252impl MechanismWorkspace {
253 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}