use ndarray::prelude::*;
use rayon::prelude::*;
use crate::{
datasets::{CatTrj, CatTrjEv, CatType, Dataset},
models::Labelled,
types::{Error, Labels, Result, Set, States},
};
#[derive(Clone, Debug)]
pub struct CatWtdTrj {
trajectory: CatTrj,
weight: f64,
}
impl TryFrom<(CatTrj, f64)> for CatWtdTrj {
type Error = Error;
fn try_from((trajectory, weight): (CatTrj, f64)) -> Result<Self> {
Self::new(trajectory, weight)
}
}
impl From<CatWtdTrj> for (CatTrj, f64) {
fn from(other: CatWtdTrj) -> Self {
(other.trajectory, other.weight)
}
}
impl CatWtdTrj {
pub fn new(trajectory: CatTrj, weight: f64) -> Result<Self> {
if !(0.0..=1.0).contains(&weight) {
return Err(Error::InvalidParameter(
"weight",
&format!("must be in the range [0, 1], but got {weight}"),
));
}
Ok(Self { trajectory, weight })
}
#[inline]
pub const fn trajectory(&self) -> &CatTrj {
&self.trajectory
}
#[inline]
pub const fn weight(&self) -> f64 {
self.weight
}
#[inline]
pub const fn states(&self) -> &States {
self.trajectory.states()
}
#[inline]
pub const fn shape(&self) -> &Array1<usize> {
self.trajectory.shape()
}
#[inline]
pub const fn times(&self) -> &Array1<f64> {
self.trajectory.times()
}
}
impl Labelled for CatWtdTrj {
#[inline]
fn labels(&self) -> &Labels {
self.trajectory.labels()
}
}
impl Dataset for CatWtdTrj {
type Values = Array2<CatType>;
type Evidence = CatTrjEv;
type EvidenceIter<'a> = <CatTrj as Dataset>::EvidenceIter<'a>;
#[inline]
fn values(&self) -> &Self::Values {
self.trajectory.values()
}
fn evidence_iter(&self) -> Self::EvidenceIter<'_> {
self.trajectory.evidence_iter()
}
#[inline]
fn sample_size(&self) -> f64 {
self.weight * (self.trajectory.values().nrows() as f64)
}
fn select(&self, x: &Set<usize>) -> Result<Self> {
let trajectory = self.trajectory.select(x)?;
let weight = self.weight;
Self::new(trajectory, weight)
}
}
#[derive(Clone, Debug)]
pub struct CatWtdTrjs {
labels: Labels,
states: States,
shape: Array1<usize>,
values: Vec<CatWtdTrj>,
}
pub struct CatWtdTrjsEvidenceIter<'a> {
trajectories: std::slice::Iter<'a, CatWtdTrj>,
current: Option<<CatWtdTrj as Dataset>::EvidenceIter<'a>>,
}
impl<'a> Iterator for CatWtdTrjsEvidenceIter<'a> {
type Item = Result<CatTrjEv>;
fn next(&mut self) -> Option<Self::Item> {
loop {
if let Some(current) = self.current.as_mut()
&& let Some(item) = current.next()
{
return Some(item);
}
self.current = self.trajectories.next().map(Dataset::evidence_iter);
self.current.as_ref()?;
}
}
}
impl CatWtdTrjs {
pub fn new<I>(values: I) -> Result<Self>
where
I: IntoIterator<Item = CatWtdTrj>,
{
let values: Vec<_> = values.into_iter().collect();
if !values
.windows(2)
.all(|trjs| trjs[0].labels().eq(trjs[1].labels()))
{
return Err(Error::IncompatibleShape("labels", "all trajectories"));
}
if !values
.windows(2)
.all(|trjs| trjs[0].states().eq(trjs[1].states()))
{
return Err(Error::IncompatibleShape("states", "all trajectories"));
}
if !values
.windows(2)
.all(|trjs| trjs[0].shape().eq(trjs[1].shape()))
{
return Err(Error::IncompatibleShape("shape", "all trajectories"));
}
let trj = values
.first()
.ok_or_else(|| Error::EmptySet("trajectories"))?;
let labels = trj.labels().clone();
let states = trj.states().clone();
let shape = trj.shape().clone();
Ok(Self {
labels,
states,
shape,
values,
})
}
#[inline]
pub fn states(&self) -> &States {
&self.states
}
#[inline]
pub fn shape(&self) -> &Array1<usize> {
&self.shape
}
}
impl FromIterator<CatWtdTrj> for CatWtdTrjs {
#[inline]
fn from_iter<I: IntoIterator<Item = CatWtdTrj>>(iter: I) -> Self {
Self::new(iter).unwrap_or_else(|e| {
log::error!("Failed to create CatWtdTrjs from iterator: {}", e);
Self {
labels: Default::default(),
states: Default::default(),
values: vec![],
shape: Array1::zeros(2),
}
})
}
}
impl FromParallelIterator<CatWtdTrj> for CatWtdTrjs {
#[inline]
fn from_par_iter<I: IntoParallelIterator<Item = CatWtdTrj>>(iter: I) -> Self {
let collected = iter.into_par_iter().collect::<Vec<_>>();
Self::new(collected).unwrap_or_else(|e| {
log::error!("Failed to create CatWtdTrjs from parallel iterator: {}", e);
Self {
labels: Default::default(),
states: Default::default(),
values: vec![],
shape: Array1::zeros(2),
}
})
}
}
impl<'a> IntoIterator for &'a CatWtdTrjs {
type IntoIter = std::slice::Iter<'a, CatWtdTrj>;
type Item = &'a CatWtdTrj;
#[inline]
fn into_iter(self) -> Self::IntoIter {
self.values.iter()
}
}
impl<'a> IntoParallelRefIterator<'a> for CatWtdTrjs {
type Item = &'a CatWtdTrj;
type Iter = rayon::slice::Iter<'a, CatWtdTrj>;
#[inline]
fn par_iter(&'a self) -> Self::Iter {
self.values.par_iter()
}
}
impl Labelled for CatWtdTrjs {
#[inline]
fn labels(&self) -> &Labels {
&self.labels
}
}
impl Dataset for CatWtdTrjs {
type Values = Vec<CatWtdTrj>;
type Evidence = CatTrjEv;
type EvidenceIter<'a> = CatWtdTrjsEvidenceIter<'a>;
#[inline]
fn values(&self) -> &Self::Values {
&self.values
}
fn evidence_iter(&self) -> Self::EvidenceIter<'_> {
CatWtdTrjsEvidenceIter {
trajectories: self.values.iter(),
current: None,
}
}
#[inline]
fn sample_size(&self) -> f64 {
self.values.iter().map(Dataset::sample_size).sum()
}
fn select(&self, x: &Set<usize>) -> Result<Self> {
Self::new(
self.values
.iter()
.map(|trj| trj.select(x))
.collect::<Result<Vec<_>>>()?,
)
}
}