use crate::ModelError;
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct FiniteScalarObservations<'a> {
values: &'a [f64],
}
impl<'a> FiniteScalarObservations<'a> {
pub fn new(values: &'a [f64]) -> Result<Self, ModelError> {
validate_scalar_observations(values)?;
Ok(Self { values })
}
#[must_use]
#[inline]
pub const fn values(&self) -> &'a [f64] {
self.values
}
}
impl<'row> ObservationView<'row> for FiniteScalarObservations<'_> {
type Observation = f64;
#[inline]
fn len(&self) -> usize {
self.values.len()
}
#[inline]
fn observation_at(&'row self, row: usize) -> Self::Observation {
self.values[row]
}
#[inline]
fn weight_at(&self, _row: usize) -> f64 {
1.0
}
#[inline]
fn validate(&self) -> Result<(), ModelError> {
validate_scalar_observations(self.values)
}
}
pub trait ObservationView<'row> {
type Observation;
fn len(&self) -> usize;
#[inline]
fn is_empty(&self) -> bool {
self.len() == 0
}
fn observation_at(&'row self, row: usize) -> Self::Observation;
fn weight_at(&self, row: usize) -> f64;
#[inline]
fn validate(&self) -> Result<(), ModelError> {
for row in 0..self.len() {
validate_observation_weight(row, self.weight_at(row))?;
}
Ok(())
}
}
impl<'row> ObservationView<'row> for &[f64] {
type Observation = f64;
#[inline]
fn len(&self) -> usize {
<[f64]>::len(self)
}
#[inline]
fn observation_at(&'row self, row: usize) -> Self::Observation {
self[row]
}
#[inline]
fn weight_at(&self, _row: usize) -> f64 {
1.0
}
#[inline]
fn validate(&self) -> Result<(), ModelError> {
Ok(())
}
}
impl<'row> ObservationView<'row> for (&[f64], &[f64]) {
type Observation = f64;
#[inline]
fn len(&self) -> usize {
self.0.len()
}
#[inline]
fn observation_at(&'row self, row: usize) -> Self::Observation {
self.0[row]
}
#[inline]
fn weight_at(&self, row: usize) -> f64 {
self.1[row]
}
fn validate(&self) -> Result<(), ModelError> {
let expected = self.0.len();
let actual = self.1.len();
if actual != expected {
return Err(ModelError::WeightLength { expected, actual });
}
for (index, weight) in self.1.iter().copied().enumerate() {
validate_observation_weight(index, weight)?;
}
Ok(())
}
}
impl<'row, const N: usize> ObservationView<'row> for &[[f64; N]] {
type Observation = [f64; N];
#[inline]
fn len(&self) -> usize {
<[[f64; N]]>::len(self)
}
#[inline]
fn observation_at(&'row self, row: usize) -> Self::Observation {
self[row]
}
#[inline]
fn weight_at(&self, _row: usize) -> f64 {
1.0
}
#[inline]
fn validate(&self) -> Result<(), ModelError> {
Ok(())
}
}
impl<'row, const N: usize> ObservationView<'row> for (&[[f64; N]], &[f64]) {
type Observation = [f64; N];
#[inline]
fn len(&self) -> usize {
self.0.len()
}
#[inline]
fn observation_at(&'row self, row: usize) -> Self::Observation {
self.0[row]
}
#[inline]
fn weight_at(&self, row: usize) -> f64 {
self.1[row]
}
fn validate(&self) -> Result<(), ModelError> {
let expected = self.0.len();
let actual = self.1.len();
if actual != expected {
return Err(ModelError::WeightLength { expected, actual });
}
for (index, weight) in self.1.iter().copied().enumerate() {
validate_observation_weight(index, weight)?;
}
Ok(())
}
}
fn validate_observation_weight(index: usize, weight: f64) -> Result<(), ModelError> {
if weight.is_finite() && weight >= 0.0 {
Ok(())
} else {
Err(ModelError::InvalidWeight { index })
}
}
fn validate_scalar_observations(values: &[f64]) -> Result<(), ModelError> {
for (index, value) in values.iter().copied().enumerate() {
if !value.is_finite() {
return Err(ModelError::InvalidObservation { index });
}
}
Ok(())
}