use ndarray::{Array1, Array3};
use crate::canonical::GaussianMixture;
use crate::graph_builder::{
assemble_model, int_attribute, make_i64_tensor, make_node, make_tensor, make_typed_value_info,
make_value_info, string_attribute, Dimension, INT64,
};
use crate::proto::{GraphProto, ModelProto};
use crate::{IR_VERSION, OPSET_VERSION};
#[must_use]
pub fn export_gaussian_mixture(model: &GaussianMixture) -> ModelProto {
let components = model.means.nrows();
let features = model.means.ncols();
let means = Array3::from_shape_fn((1, components, features), |(_, component, feature)| {
model.means[(component, feature)] as f32
});
let precisions = model.precisions.mapv(|value| value as f32);
let offsets = model.offsets.mapv(|value| value as f32);
let minus_half = Array1::from_vec(vec![-0.5_f32]);
let graph = GraphProto {
node: vec![
make_node("Unsqueeze", ["X", "axis_one"], ["expanded"], Vec::new()),
make_node("Sub", ["expanded", "means"], ["delta"], Vec::new()),
make_node(
"Einsum",
["delta", "precisions", "delta"],
["mahalanobis"],
vec![string_attribute("equation", b"bcf,cfg,bcg->bc".to_vec())],
),
make_node(
"Mul",
["mahalanobis", "minus_half"],
["quadratic"],
Vec::new(),
),
make_node("Add", ["quadratic", "offsets"], ["scores"], Vec::new()),
make_node(
"ArgMax",
["scores"],
["component"],
vec![int_attribute("axis", 1), int_attribute("keepdims", 0)],
),
],
name: "gaussian_mixture".into(),
initializer: vec![
make_tensor("means", &means.into_dyn()),
make_tensor("precisions", &precisions.into_dyn()),
make_tensor("offsets", &offsets.into_dyn()),
make_tensor("minus_half", &minus_half.into_dyn()),
make_i64_tensor("axis_one", &[1], vec![1]),
],
doc_string: String::new(),
input: vec![make_value_info(
"X",
&[
Dimension::Symbolic("batch".into()),
Dimension::Fixed(features),
],
)],
output: vec![make_typed_value_info(
"component",
&[Dimension::Symbolic("batch".into())],
INT64,
)],
value_info: Vec::new(),
};
assemble_model(graph, OPSET_VERSION, IR_VERSION)
}