use ndarray::prelude::*;
use crate::{
datasets::{CatEv, CatSample, CatTable, Dataset},
models::Labelled,
types::{Error, Labels, Result, Set, States},
};
pub type CatWtdSample = (CatSample, f64);
#[derive(Clone, Debug)]
pub struct CatWtdTable {
dataset: CatTable,
weights: Array1<f64>,
}
impl Labelled for CatWtdTable {
#[inline]
fn labels(&self) -> &Labels {
self.dataset.labels()
}
}
impl CatWtdTable {
pub fn new(dataset: CatTable, weights: Array1<f64>) -> Result<Self> {
if dataset.values().nrows() != weights.len() {
return Err(Error::InvalidParameter(
"weights",
"must have the same length as the dataset",
));
}
if !weights.iter().all(|&w| w.is_finite()) {
return Err(Error::InvalidParameter("weights", "must be finite"));
}
Ok(Self { dataset, weights })
}
#[inline]
pub const fn states(&self) -> &States {
self.dataset.states()
}
#[inline]
pub const fn shape(&self) -> &Array1<usize> {
self.dataset.shape()
}
#[inline]
pub const fn weights(&self) -> &Array1<f64> {
&self.weights
}
}
impl Dataset for CatWtdTable {
type Values = CatTable;
type Evidence = CatEv;
type EvidenceIter<'a> = <CatTable as Dataset>::EvidenceIter<'a>;
#[inline]
fn values(&self) -> &Self::Values {
&self.dataset
}
fn evidence_iter(&self) -> Self::EvidenceIter<'_> {
self.dataset.evidence_iter()
}
#[inline]
fn sample_size(&self) -> f64 {
self.weights.sum()
}
fn select(&self, x: &Set<usize>) -> Result<Self> {
let dataset = self.dataset.select(x)?;
let weights = self.weights.clone();
Self::new(dataset, weights)
}
}
impl From<CatTable> for CatWtdTable {
#[inline]
fn from(dataset: CatTable) -> Self {
let weights = Array::ones(dataset.values().nrows());
Self { dataset, weights }
}
}