#![allow(clippy::cast_possible_truncation)]
use std::sync::Arc;
use crate::error::ModelError;
#[derive(Clone, Debug, Default)]
pub struct ValueBatch {
pub n_rows: usize,
pub n_nodes: usize,
pub values: Arc<[f64]>,
}
impl ValueBatch {
#[must_use]
pub fn zeros(n_rows: usize, n_nodes: usize) -> Self {
Self { n_rows, n_nodes, values: Arc::from(vec![0.0; n_rows.saturating_mul(n_nodes)]) }
}
pub fn column(&self, node: usize) -> Result<&[f64], ModelError> {
if node >= self.n_nodes {
return Err(ModelError::Shape { message: "value column out of range".into() });
}
let start = node * self.n_rows;
Ok(&self.values[start..start + self.n_rows])
}
pub fn get(&self, row: usize, node: usize) -> Result<f64, ModelError> {
if row >= self.n_rows || node >= self.n_nodes {
return Err(ModelError::Shape { message: "value index out of range".into() });
}
Ok(self.values[node * self.n_rows + row])
}
}
#[derive(Debug)]
pub struct ValueBatchMut<'a> {
pub n_rows: usize,
pub n_nodes: usize,
pub values: &'a mut [f64],
}
impl<'a> ValueBatchMut<'a> {
pub fn new(n_rows: usize, n_nodes: usize, values: &'a mut [f64]) -> Result<Self, ModelError> {
if values.len() < n_rows.saturating_mul(n_nodes) {
return Err(ModelError::Shape { message: "value buffer too short".into() });
}
Ok(Self { n_rows, n_nodes, values })
}
pub fn column_mut(&mut self, node: usize) -> Result<&mut [f64], ModelError> {
if node >= self.n_nodes {
return Err(ModelError::Shape { message: "value column out of range".into() });
}
let start = node * self.n_rows;
Ok(&mut self.values[start..start + self.n_rows])
}
pub fn set(&mut self, row: usize, node: usize, v: f64) -> Result<(), ModelError> {
if row >= self.n_rows || node >= self.n_nodes {
return Err(ModelError::Shape { message: "value index out of range".into() });
}
self.values[node * self.n_rows + row] = v;
Ok(())
}
#[must_use]
pub fn into_batch(self) -> ValueBatch {
ValueBatch {
n_rows: self.n_rows,
n_nodes: self.n_nodes,
values: Arc::from(self.values.to_vec()),
}
}
}
#[derive(Clone, Debug, Default)]
pub struct NoiseBatch {
pub n_rows: usize,
pub n_nodes: usize,
pub values: Arc<[f64]>,
}
impl NoiseBatch {
#[must_use]
pub fn zeros(n_rows: usize, n_nodes: usize) -> Self {
Self { n_rows, n_nodes, values: Arc::from(vec![0.0; n_rows.saturating_mul(n_nodes)]) }
}
pub fn column(&self, node: usize) -> Result<&[f64], ModelError> {
if node >= self.n_nodes {
return Err(ModelError::Shape { message: "noise column out of range".into() });
}
let start = node * self.n_rows;
Ok(&self.values[start..start + self.n_rows])
}
}
#[derive(Debug)]
pub struct NoiseBatchMut<'a> {
pub n_rows: usize,
pub n_nodes: usize,
pub values: &'a mut [f64],
}
impl<'a> NoiseBatchMut<'a> {
pub fn new(n_rows: usize, n_nodes: usize, values: &'a mut [f64]) -> Result<Self, ModelError> {
if values.len() < n_rows.saturating_mul(n_nodes) {
return Err(ModelError::Shape { message: "noise buffer too short".into() });
}
Ok(Self { n_rows, n_nodes, values })
}
pub fn column(&self, node: usize) -> Result<&[f64], ModelError> {
if node >= self.n_nodes {
return Err(ModelError::Shape { message: "noise column out of range".into() });
}
let start = node * self.n_rows;
Ok(&self.values[start..start + self.n_rows])
}
pub fn column_mut(&mut self, node: usize) -> Result<&mut [f64], ModelError> {
if node >= self.n_nodes {
return Err(ModelError::Shape { message: "noise column out of range".into() });
}
let start = node * self.n_rows;
Ok(&mut self.values[start..start + self.n_rows])
}
#[must_use]
pub fn into_batch(self) -> NoiseBatch {
NoiseBatch {
n_rows: self.n_rows,
n_nodes: self.n_nodes,
values: Arc::from(self.values.to_vec()),
}
}
}
#[derive(Clone, Copy, Debug)]
pub struct ParentBatch<'a> {
pub n_rows: usize,
pub n_parents: usize,
pub values: &'a [f64],
}
impl<'a> ParentBatch<'a> {
#[must_use]
pub const fn empty(n_rows: usize) -> Self {
Self { n_rows, n_parents: 0, values: &[] }
}
pub fn column(&self, p: usize) -> Result<&'a [f64], ModelError> {
if p >= self.n_parents {
return Err(ModelError::Shape { message: "parent column out of range".into() });
}
let start = p * self.n_rows;
Ok(&self.values[start..start + self.n_rows])
}
}
#[derive(Clone, Debug, Default)]
pub struct MechanismWorkspace {
pub parents: Vec<f64>,
pub scratch: Vec<f64>,
pub grow_count: u32,
}
impl MechanismWorkspace {
pub fn prepare(&mut self, n_rows: usize, n_parents: usize) {
let need = n_rows.saturating_mul(n_parents.max(1));
if self.parents.capacity() < need {
self.parents.reserve(need - self.parents.capacity());
self.grow_count = self.grow_count.saturating_add(1);
}
self.parents.resize(need, 0.0);
if self.scratch.capacity() < n_rows {
self.scratch.reserve(n_rows.saturating_sub(self.scratch.capacity()));
self.grow_count = self.grow_count.saturating_add(1);
}
self.scratch.resize(n_rows, 0.0);
}
}