use crate::classification::{fclassif_knn_fit, fclassif_lda_fit, ClassifFit};
use crate::error::FdarError;
use crate::explain_generic::{FpcPredictor, TaskType};
use crate::matrix::FdMatrix;
use crate::shapelet::discovery::{ShapeletDiscoveryConfig, ShapeletSet};
use crate::shapelet::transform::{shapelet_transform_fit, ShapeletTransformFit};
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[non_exhaustive]
pub enum ShapeletClassifier {
Knn {
k: usize,
},
Lda,
}
impl Default for ShapeletClassifier {
fn default() -> Self {
Self::Knn { k: 1 }
}
}
#[derive(Debug, Clone, PartialEq, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct ShapeletClassifierConfig {
pub discovery: ShapeletDiscoveryConfig,
pub classifier: ShapeletClassifier,
pub ncomp: Option<usize>,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[non_exhaustive]
pub struct ShapeletClassifierFit {
pub transform: ShapeletTransformFit,
pub classifier: ClassifFit,
pub config: ShapeletClassifierConfig,
pub classes: Vec<usize>,
}
impl ShapeletClassifierFit {
#[must_use]
pub fn shapelets(&self) -> &ShapeletSet {
self.transform.shapelets()
}
#[must_use]
pub fn transform(&self) -> &ShapeletTransformFit {
&self.transform
}
#[must_use]
pub fn classifier(&self) -> &ClassifFit {
&self.classifier
}
#[must_use]
pub fn train_accuracy(&self) -> f64 {
self.classifier.result.accuracy
}
#[must_use = "predicted labels should not be discarded"]
pub fn predict(&self, new_data: &FdMatrix) -> Result<Vec<usize>, FdarError> {
let features = self.transform.transform(new_data)?;
let scores = self.classifier.project(&features);
let d = scores.ncols();
let n_new = scores.nrows();
let task = self.classifier.task_type();
let mut out = Vec::with_capacity(n_new);
for i in 0..n_new {
let row: Vec<f64> = (0..d).map(|j| scores[(i, j)]).collect();
let raw = self.classifier.predict_from_scores(&row, None);
let remapped = match task {
TaskType::BinaryClassification => usize::from(raw >= 0.5),
TaskType::MulticlassClassification(_) => raw.round() as usize,
TaskType::Regression => raw.round() as usize,
};
let label = self.classes.get(remapped).copied().unwrap_or(remapped);
out.push(label);
}
Ok(out)
}
}
#[must_use = "the fitted classifier should not be discarded"]
pub fn shapelet_classifier_fit(
data: &FdMatrix,
labels: &[usize],
config: &ShapeletClassifierConfig,
) -> Result<ShapeletClassifierFit, FdarError> {
let transform = shapelet_transform_fit(data, labels, &config.discovery)?;
let features = transform.features().clone();
let k = transform.shapelets().len();
let n = features.nrows();
let ncomp = config
.ncomp
.unwrap_or(k)
.min(k)
.min(n.saturating_sub(1))
.max(1);
let classifier = match config.classifier {
ShapeletClassifier::Knn { k: k_nn } => {
fclassif_knn_fit(&features, labels, None, ncomp, k_nn)?
}
ShapeletClassifier::Lda => fclassif_lda_fit(&features, labels, None, ncomp)?,
};
let mut classes: Vec<usize> = labels.to_vec();
classes.sort_unstable();
classes.dedup();
Ok(ShapeletClassifierFit {
transform,
classifier,
config: config.clone(),
classes,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::shapelet::discovery::ShapeletDiscoveryConfig;
fn labeled_dataset(n: usize, m: usize) -> (FdMatrix, Vec<usize>) {
let mut flat = vec![0.0f64; n * m];
let mut labels = vec![0usize; n];
let motif_start = m / 2;
let motif_len = (m / 4).max(1);
for i in 0..n {
let class1 = i % 2 == 1;
labels[i] = usize::from(class1);
let offset = 0.01 * (i as f64);
for j in 0..m {
let mut v = offset + (j as f64) * 0.001;
let hash = (i.wrapping_mul(2654435761) ^ j.wrapping_mul(40503)) % 211;
v += 0.05 * (hash as f64 / 211.0 - 0.5);
if class1 && j >= motif_start && j < motif_start + motif_len {
let k = j - motif_start;
let half = motif_len / 2;
let tri = if k <= half {
k as f64
} else {
(motif_len - k) as f64
};
v += tri;
}
flat[i + j * n] = v;
}
}
(FdMatrix::from_column_major(flat, n, m).unwrap(), labels)
}
fn discovery_cfg() -> ShapeletDiscoveryConfig {
ShapeletDiscoveryConfig {
min_length: 3,
max_length: 6,
max_candidates: None,
max_shapelets: 4,
seed: 0,
..Default::default()
}
}
#[test]
fn test_stc_fit_predict_end_to_end() {
let (train, train_y) = labeled_dataset(24, 24);
let (test, test_y) = labeled_dataset(12, 24);
let cfg = ShapeletClassifierConfig {
discovery: discovery_cfg(),
..Default::default()
};
let fit = shapelet_classifier_fit(&train, &train_y, &cfg).unwrap();
let preds = fit.predict(&test).unwrap();
assert_eq!(preds.len(), test_y.len());
let correct = preds.iter().zip(&test_y).filter(|(p, t)| p == t).count();
let acc = correct as f64 / test_y.len() as f64;
assert!(
acc > 0.6,
"held-out accuracy {acc} should be well above chance (0.5)"
);
}
#[test]
fn test_stc_knn_default() {
assert_eq!(
ShapeletClassifierConfig::default().classifier,
ShapeletClassifier::Knn { k: 1 }
);
let (train, train_y) = labeled_dataset(20, 24);
let cfg = ShapeletClassifierConfig {
discovery: discovery_cfg(),
..Default::default()
};
let fit = shapelet_classifier_fit(&train, &train_y, &cfg).unwrap();
let acc = fit.train_accuracy();
assert!(
(0.0..=1.0).contains(&acc),
"train_accuracy out of range: {acc}"
);
}
#[test]
fn test_stc_lda_option() {
let (train, train_y) = labeled_dataset(24, 24);
let (test, _test_y) = labeled_dataset(10, 24);
let cfg = ShapeletClassifierConfig {
discovery: discovery_cfg(),
classifier: ShapeletClassifier::Lda,
ncomp: None,
};
let fit = shapelet_classifier_fit(&train, &train_y, &cfg).unwrap();
let preds = fit.predict(&test).unwrap();
assert_eq!(preds.len(), 10);
for &p in &preds {
assert!(p == 0 || p == 1, "unexpected label {p}");
}
}
#[test]
fn test_stc_predict_consistency() {
let (train, train_y) = labeled_dataset(24, 24);
let cfg = ShapeletClassifierConfig {
discovery: discovery_cfg(),
classifier: ShapeletClassifier::Lda,
ncomp: None,
};
let fit = shapelet_classifier_fit(&train, &train_y, &cfg).unwrap();
let fit_time: Vec<usize> = fit
.classifier
.result
.predicted
.iter()
.map(|&r| fit.classes[r])
.collect();
let re = fit.predict(&train).unwrap();
assert_eq!(
re, fit_time,
"predict(train) != fit-time training predictions"
);
}
#[test]
fn test_stc_validation() {
let (data, _labels) = labeled_dataset(8, 24);
let single = vec![0usize; 8];
let cfg = ShapeletClassifierConfig {
discovery: discovery_cfg(),
..Default::default()
};
assert!(shapelet_classifier_fit(&data, &single, &cfg).is_err());
let short = vec![0usize, 1, 0];
assert!(shapelet_classifier_fit(&data, &short, &cfg).is_err());
}
#[test]
fn test_shapelet_reexports() {
use crate::{
discover_shapelets, shapelet_classifier_fit as _scf, shapelet_distance,
shapelet_transform, shapelet_transform_fit, QualityMeasure, Shapelet,
ShapeletClassifier, ShapeletClassifierConfig, ShapeletClassifierFit,
ShapeletDiscoveryConfig, ShapeletSet, ShapeletTransformFit,
};
let _ = _scf;
let _ = shapelet_distance;
let _ = discover_shapelets;
let _ = shapelet_transform;
let _ = shapelet_transform_fit;
let _c: fn() -> ShapeletClassifierConfig = ShapeletClassifierConfig::default;
let _q = QualityMeasure::InfoGain;
let _cl = ShapeletClassifier::default();
fn _takes(
_a: &Shapelet,
_b: &ShapeletSet,
_c: &ShapeletDiscoveryConfig,
_d: &ShapeletTransformFit,
_e: &ShapeletClassifierFit,
) {
}
}
}