use crate::canonical::CategoricalNaiveBayes;
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_categorical_naive_bayes(model: &CategoricalNaiveBayes) -> ModelProto {
let features = model.feature_log_probabilities.len();
let classes = model.class_labels.len();
let mut nodes = vec![make_node(
"Cast",
["X"],
["categories"],
vec![int_attribute("to", INT64 as i64)],
)];
let mut initializer = vec![
make_tensor(
"log_priors",
&model.log_priors.mapv(|value| value as f32).into_dyn(),
),
make_i64_tensor("labels", &[classes], model.class_labels.clone()),
];
let mut sum = "log_priors".to_owned();
for (feature, table) in model.feature_log_probabilities.iter().enumerate() {
let index_name = format!("feature_index_{feature}");
let table_name = format!("table_{feature}");
let category_name = format!("category_{feature}");
let score_name = format!("feature_score_{feature}");
let sum_name = format!("score_sum_{feature}");
initializer.push(make_i64_tensor(&index_name, &[], vec![feature as i64]));
initializer.push(make_tensor(
&table_name,
&table.mapv(|value| value as f32).into_dyn(),
));
nodes.extend([
make_node(
"Gather",
["categories", &index_name],
[&category_name],
vec![int_attribute("axis", 1)],
),
make_node(
"Gather",
[&table_name, &category_name],
[&score_name],
vec![int_attribute("axis", 0)],
),
make_node("Add", [&sum, &score_name], [&sum_name], Vec::new()),
]);
sum = sum_name;
}
nodes.extend([
make_node(
"ArgMax",
[&sum],
["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: "categorical_naive_bayes".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::new(),
},
OPSET_VERSION,
IR_VERSION,
)
}