use itertools::Itertools;
use ndarray::prelude::*;
use rayon::prelude::*;
use crate::{
datasets::{CatTable, CatTrjEv, CatTrjEvT, CatType, Dataset},
models::Labelled,
types::{Error, Labels, Result, Set, States},
};
#[derive(Clone, Debug)]
pub struct CatTrj {
events: CatTable,
times: Array1<f64>,
}
pub struct CatTrjEvidenceIter<'a> {
rows: ndarray::iter::LanesIter<'a, CatType, Ix1>,
time_bounds: std::vec::IntoIter<(f64, f64)>,
states: &'a States,
}
impl<'a> Iterator for CatTrjEvidenceIter<'a> {
type Item = Result<CatTrjEv>;
fn next(&mut self) -> Option<Self::Item> {
let row = self.rows.next()?;
let (start_time, end_time) = self.time_bounds.next().unwrap_or((0.0, 0.0));
let evidences =
row.iter()
.enumerate()
.map(|(event, &state)| CatTrjEvT::CertainPositiveInterval {
event,
state: state as usize,
start_time,
end_time,
});
Some(CatTrjEv::new(self.states.clone(), evidences))
}
}
impl CatTrj {
pub fn new(
states: States,
mut events: Array2<CatType>,
mut times: Array1<f64>,
) -> Result<Self> {
if events.nrows() != times.len() {
return Err(Error::IncompatibleShape(
&events.nrows().to_string(),
×.len().to_string(),
));
}
times.iter().try_for_each(|&t| {
if !t.is_finite() || t < 0. {
return Err(Error::InvalidParameter(
"times",
&format!("value must be finite and positive, found {t}"),
));
}
Ok(())
})?;
let mut sorted_idx: Vec<_> = (0..events.nrows()).collect();
sorted_idx.sort_by(|&a, &b| {
times[a]
.partial_cmp(×[b])
.unwrap_or_else(|| unreachable!())
});
if !sorted_idx.iter().is_sorted() {
let mut new_times = times.clone();
new_times
.iter_mut()
.enumerate()
.for_each(|(i, new_time)| *new_time = times[sorted_idx[i]]);
times = new_times;
let mut new_events = events.clone();
new_events
.rows_mut()
.into_iter()
.enumerate()
.for_each(|(i, mut new_events_row)| {
new_events_row.assign(&events.row(sorted_idx[i]));
});
events = new_events;
}
{
let count = times.iter().dedup().count();
let length = times.len();
if count != length {
return Err(Error::InvalidParameter(
"times",
&format!("must be unique, found {} duplicates", length - count),
));
}
}
for ((e_i, _), (e_j, _)) in events.rows().into_iter().zip(×).tuple_windows() {
let count = e_i.iter().zip(e_j).filter(|(a, b)| a != b).count();
if count > 1 {
return Err(Error::InvalidParameter(
"events",
&format!("must contain at max one change per transition, found {count}"),
));
}
}
let events = CatTable::new(states, events)?;
Ok(Self { events, times })
}
#[inline]
pub const fn states(&self) -> &States {
self.events.states()
}
#[inline]
pub const fn shape(&self) -> &Array1<usize> {
self.events.shape()
}
#[inline]
pub const fn times(&self) -> &Array1<f64> {
&self.times
}
}
impl Labelled for CatTrj {
#[inline]
fn labels(&self) -> &Labels {
self.events.labels()
}
}
impl Dataset for CatTrj {
type Values = Array2<CatType>;
type Evidence = CatTrjEv;
type EvidenceIter<'a> = CatTrjEvidenceIter<'a>;
#[inline]
fn values(&self) -> &Self::Values {
self.events.values()
}
fn evidence_iter(&self) -> Self::EvidenceIter<'_> {
let mut end_times: Vec<f64> = self.times.iter().copied().skip(1).collect();
end_times.push(*self.times.last().unwrap_or(&0.0));
let time_bounds: Vec<(f64, f64)> = self.times.iter().copied().zip(end_times).collect();
CatTrjEvidenceIter {
rows: self.values().rows().into_iter(),
time_bounds: time_bounds.into_iter(),
states: self.states(),
}
}
#[inline]
fn sample_size(&self) -> f64 {
self.events.values().nrows() as f64
}
fn select(&self, x: &Set<usize>) -> Result<Self> {
let events = self.events.select(x)?;
let states = events.states().clone();
let events = events.values().clone();
let times = self.times.clone();
Self::new(states, events, times)
}
}
#[derive(Clone, Debug)]
pub struct CatTrjs {
labels: Labels,
states: States,
shape: Array1<usize>,
values: Vec<CatTrj>,
}
impl CatTrjs {
pub fn new<I>(values: I) -> Result<Self>
where
I: IntoIterator<Item = CatTrj>,
{
let values: Vec<_> = values.into_iter().collect();
if !values
.windows(2)
.all(|trjs| trjs[0].labels().eq(trjs[1].labels()))
{
return Err(Error::ConstructionError(
"All trajectories must have the same labels.",
));
}
if !values
.windows(2)
.all(|trjs| trjs[0].states().eq(trjs[1].states()))
{
return Err(Error::ConstructionError(
"All trajectories must have the same states.",
));
}
if !values
.windows(2)
.all(|trjs| trjs[0].shape().eq(trjs[1].shape()))
{
return Err(Error::ConstructionError(
"All trajectories must have the same shape.",
));
}
let (labels, states, shape) = match values.first() {
None => (Labels::default(), States::default(), Array1::default((0,))),
Some(x) => (x.labels().clone(), x.states().clone(), x.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<CatTrj> for CatTrjs {
#[inline]
fn from_iter<I: IntoIterator<Item = CatTrj>>(iter: I) -> Self {
Self::new(iter).unwrap_or_else(|e| {
log::error!("Failed to create CatTrjs from iterator: {}", e);
Self {
labels: Default::default(),
states: Default::default(),
values: vec![],
shape: Array1::zeros(2),
}
})
}
}
impl FromParallelIterator<CatTrj> for CatTrjs {
#[inline]
fn from_par_iter<I: IntoParallelIterator<Item = CatTrj>>(iter: I) -> Self {
let collected = iter.into_par_iter().collect::<Vec<_>>();
Self::new(collected).unwrap_or_else(|e| {
log::error!("Failed to create CatTrjs from parallel iterator: {}", e);
Self {
labels: Default::default(),
states: Default::default(),
values: vec![],
shape: Array1::zeros(2),
}
})
}
}
impl<'a> IntoIterator for &'a CatTrjs {
type IntoIter = std::slice::Iter<'a, CatTrj>;
type Item = &'a CatTrj;
#[inline]
fn into_iter(self) -> Self::IntoIter {
self.values.iter()
}
}
impl<'a> IntoParallelRefIterator<'a> for CatTrjs {
type Item = &'a CatTrj;
type Iter = rayon::slice::Iter<'a, CatTrj>;
#[inline]
fn par_iter(&'a self) -> Self::Iter {
self.values.par_iter()
}
}
impl Labelled for CatTrjs {
#[inline]
fn labels(&self) -> &Labels {
&self.labels
}
}
pub struct CatTrjsEvidenceIter<'a> {
trajectories: std::slice::Iter<'a, CatTrj>,
current: Option<<CatTrj as Dataset>::EvidenceIter<'a>>,
}
impl<'a> Iterator for CatTrjsEvidenceIter<'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 Dataset for CatTrjs {
type Values = Vec<CatTrj>;
type Evidence = CatTrjEv;
type EvidenceIter<'a> = CatTrjsEvidenceIter<'a>;
#[inline]
fn values(&self) -> &Self::Values {
&self.values
}
fn evidence_iter(&self) -> Self::EvidenceIter<'_> {
CatTrjsEvidenceIter {
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<_>>>()?,
)
}
}