use ndarray::{Array1, Array2, Array3};
use crate::canonical::DbscanModel;
use crate::graph_builder::{
assemble_model, int_attribute, make_i64_tensor, make_node, make_tensor, make_typed_value_info,
make_value_info, Dimension, INT64,
};
use crate::proto::{GraphProto, ModelProto};
use crate::{IR_VERSION, OPSET_VERSION};
#[must_use]
pub fn export_dbscan(model: &DbscanModel) -> ModelProto {
let samples = model.samples.nrows();
let features = model.samples.ncols();
let classes_with_noise = model.n_clusters + 1;
let sample_tensor = Array3::from_shape_fn((1, samples, features), |(_, row, column)| {
model.samples[(row, column)] as f32
});
let votes = Array2::from_shape_fn((samples, classes_with_noise), |(sample, class)| {
let fitted = model.cluster_indices[sample];
let fitted_class = if fitted < 0 {
model.n_clusters
} else {
fitted as usize
};
if class == fitted_class {
1.0
} else {
0.0
}
});
let graph = GraphProto {
node: vec![
make_node("Unsqueeze", ["X", "axis_one"], ["expanded"], Vec::new()),
make_node("Sub", ["expanded", "samples"], ["delta"], Vec::new()),
make_node("Mul", ["delta", "delta"], ["squared"], Vec::new()),
make_node(
"ReduceSum",
["squared", "axis_features"],
["squared_distances"],
vec![int_attribute("keepdims", 0)],
),
make_node(
"LessOrEqual",
["squared_distances", "epsilon_squared"],
["neighbor_mask"],
Vec::new(),
),
make_node(
"Cast",
["neighbor_mask"],
["neighbors"],
vec![int_attribute("to", 1)],
),
make_node("MatMul", ["neighbors", "votes"], ["counts"], Vec::new()),
make_node(
"ArgMax",
["counts"],
["class_index"],
vec![int_attribute("axis", 1), int_attribute("keepdims", 0)],
),
make_node(
"Equal",
["class_index", "noise_index"],
["is_noise"],
Vec::new(),
),
make_node("Add", ["class_index", "one"], ["cluster_label"], Vec::new()),
make_node(
"Where",
["is_noise", "zero", "cluster_label"],
["label"],
Vec::new(),
),
],
name: "dbscan".into(),
initializer: vec![
make_tensor("samples", &sample_tensor.into_dyn()),
make_tensor("votes", &votes.into_dyn()),
make_tensor(
"epsilon_squared",
&Array1::from_vec(vec![(model.epsilon * model.epsilon) as f32]).into_dyn(),
),
make_i64_tensor("axis_one", &[1], vec![1]),
make_i64_tensor("axis_features", &[1], vec![2]),
make_i64_tensor("noise_index", &[1], vec![model.n_clusters as i64]),
make_i64_tensor("one", &[1], vec![1]),
make_i64_tensor("zero", &[1], vec![0]),
],
doc_string: String::new(),
input: vec![make_value_info(
"X",
&[
Dimension::Symbolic("batch".into()),
Dimension::Fixed(features),
],
)],
output: vec![make_typed_value_info(
"label",
&[Dimension::Symbolic("batch".into())],
INT64,
)],
value_info: Vec::new(),
};
assemble_model(graph, OPSET_VERSION, IR_VERSION)
}