onnx-export-rs 0.1.1

Export canonical Rust machine-learning models to ONNX
Documentation
use ndarray::{Array1, Array2};

use crate::{Error, Result};

/// Fitted Gaussian Naive Bayes inference parameters.
#[derive(Clone, Debug, PartialEq)]
pub struct GaussianNaiveBayes {
    /// Per-class, per-feature means shaped `[classes, features]`.
    pub means: Array2<f64>,
    /// Per-class, per-feature variances shaped `[classes, features]`.
    pub variances: Array2<f64>,
    /// Prior probability for each class.
    pub priors: Array1<f64>,
    /// Integer label corresponding to each class row.
    pub class_labels: Vec<i64>,
}

impl GaussianNaiveBayes {
    /// Creates validated Gaussian Naive Bayes parameters.
    ///
    /// # Errors
    ///
    /// Returns an error when shapes disagree, variances/priors are not
    /// positive, labels are duplicated, or a numeric value is non-finite.
    pub fn new(
        means: Array2<f64>,
        variances: Array2<f64>,
        priors: Array1<f64>,
        class_labels: Vec<i64>,
    ) -> Result<Self> {
        let classes = means.nrows();
        let mut unique_labels = class_labels.clone();
        unique_labels.sort_unstable();
        unique_labels.dedup();
        if classes == 0
            || means.ncols() == 0
            || variances.dim() != means.dim()
            || priors.len() != classes
            || class_labels.len() != classes
            || unique_labels.len() != classes
            || means.iter().any(|value| !value.is_finite())
            || variances
                .iter()
                .any(|value| !value.is_finite() || *value <= 0.0)
            || priors
                .iter()
                .any(|value| !value.is_finite() || *value <= 0.0)
        {
            return Err(Error::InvalidModel(
                "invalid Gaussian Naive Bayes parameters".into(),
            ));
        }
        Ok(Self {
            means,
            variances,
            priors,
            class_labels,
        })
    }

    /// Required number of input features.
    #[must_use]
    pub fn n_features(&self) -> usize {
        self.means.ncols()
    }
}