use itertools::Itertools;
use ndarray::prelude::*;
use rayon::prelude::*;
use crate::{
datasets::{CatTable, CatType, Dataset},
models::Labelled,
types::{Labels, Set, States},
};
#[derive(Clone, Debug)]
pub struct CatTrj {
events: CatTable,
times: Array1<f64>,
}
impl CatTrj {
pub fn new(states: States, mut events: Array2<CatType>, mut times: Array1<f64>) -> Self {
assert_eq!(
events.nrows(),
times.len(),
"Trajectory events and times must have the same length."
);
times.iter().for_each(|&t| {
assert!(
t.is_finite() && t >= 0.,
"Trajectory times must be finite and positive: \n\
\t expected: time >= 0 , \n\
\t found: time == {t} ."
);
});
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();
assert_eq!(
count, length,
"Trajectory times must be unique: \n\
\t expected: {count} deduplicated time-points, \n\
\t found: {length} non-deduplicated time-points, \n\
\t for: {times}."
);
}
events
.rows()
.into_iter()
.zip(×)
.tuple_windows()
.for_each(|((e_i, t_i), (e_j, t_j))| {
let count = e_i.iter().zip(e_j).filter(|(a, b)| a != b).count();
assert!(
count <= 1,
"Trajectory events must contain at max one change per transition: \n\
\t expected: count <= 1 state change, \n\
\t found: count == {count} state changes, \n\
\t for: {e_i} event with time {t_i}, \n\
\t and: {e_j} event with time {t_j}."
);
});
let events = CatTable::new(states, events);
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>;
#[inline]
fn values(&self) -> &Self::Values {
self.events.values()
}
#[inline]
fn sample_size(&self) -> f64 {
self.events.values().nrows() as f64
}
fn select(&self, x: &Set<usize>) -> 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) -> Self
where
I: IntoIterator<Item = CatTrj>,
{
let values: Vec<_> = values.into_iter().collect();
assert!(
values
.windows(2)
.all(|trjs| trjs[0].labels().eq(trjs[1].labels())),
"All trajectories must have the same labels."
);
assert!(
values
.windows(2)
.all(|trjs| trjs[0].states().eq(trjs[1].states())),
"All trajectories must have the same states."
);
assert!(
values
.windows(2)
.all(|trjs| trjs[0].shape().eq(trjs[1].shape())),
"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()),
};
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)
}
}
impl FromParallelIterator<CatTrj> for CatTrjs {
#[inline]
fn from_par_iter<I: IntoParallelIterator<Item = CatTrj>>(iter: I) -> Self {
Self::new(iter.into_par_iter().collect::<Vec<_>>())
}
}
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
}
}
impl Dataset for CatTrjs {
type Values = Vec<CatTrj>;
#[inline]
fn values(&self) -> &Self::Values {
&self.values
}
#[inline]
fn sample_size(&self) -> f64 {
self.values.iter().map(Dataset::sample_size).sum()
}
fn select(&self, x: &Set<usize>) -> Self {
Self::new(self.values.iter().map(|trj| trj.select(x)))
}
}