use ndarray::Array1;
use crate::canonical::LinearScoreClassifier;
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_linear_score_classifier(model: &LinearScoreClassifier) -> ModelProto {
let features = model.coefficients.nrows();
let classes = model.coefficients.ncols();
let coefficients = model.coefficients.mapv(|value| value as f32);
let bias = model.bias.mapv(|value| value as f32);
let mut initializer = vec![
make_tensor("coefficients", &coefficients.into_dyn()),
make_tensor("bias", &bias.into_dyn()),
make_i64_tensor("labels", &[classes], model.class_labels.clone()),
];
let mut nodes = Vec::new();
let score_input = if let Some(threshold) = model.binarize {
initializer.push(make_tensor(
"threshold",
&Array1::from_vec(vec![threshold as f32]).into_dyn(),
));
nodes.extend([
make_node("Greater", ["X", "threshold"], ["binary_bool"], Vec::new()),
make_node(
"Cast",
["binary_bool"],
["binary"],
vec![int_attribute("to", 1)],
),
]);
"binary"
} else {
"X"
};
nodes.extend([
make_node(
"Gemm",
[score_input, "coefficients", "bias"],
["scores"],
Vec::new(),
),
make_node(
"ArgMax",
["scores"],
["class_index"],
vec![int_attribute("axis", 1), int_attribute("keepdims", 0)],
),
make_node(
"Gather",
["labels", "class_index"],
["label"],
vec![int_attribute("axis", 0)],
),
]);
assemble_model(
GraphProto {
node: nodes,
name: "linear_score_classifier".into(),
initializer,
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![make_value_info(
"scores",
&[
Dimension::Symbolic("batch".into()),
Dimension::Fixed(classes),
],
)],
},
OPSET_VERSION,
IR_VERSION,
)
}