use ndarray::{Array1, Array2};
use crate::canonical::{GeneralizedLinearModel, LinkFunction};
use crate::graph_builder::{assemble_model, make_node, make_tensor, make_value_info, Dimension};
use crate::proto::{GraphProto, ModelProto};
use crate::{IR_VERSION, OPSET_VERSION};
#[must_use]
pub fn export_generalized_linear(model: &GeneralizedLinearModel) -> ModelProto {
let features = model.n_features();
let coefficients = Array2::from_shape_fn((features, 1), |(feature, _)| {
model.coefficients[feature] as f32
});
let intercept = Array1::from_vec(vec![model.intercept as f32]);
let (gemm_output, mut nodes) = match model.link {
LinkFunction::Identity => ("Y", Vec::new()),
LinkFunction::Log => (
"scores",
vec![make_node("Exp", ["scores"], ["Y"], Vec::new())],
),
LinkFunction::Logit => (
"scores",
vec![make_node("Sigmoid", ["scores"], ["Y"], Vec::new())],
),
};
nodes.insert(
0,
make_node("Gemm", ["X", "W", "b"], [gemm_output], Vec::new()),
);
let graph = GraphProto {
node: nodes,
name: "generalized_linear_model".into(),
initializer: vec![
make_tensor("W", &coefficients.into_dyn()),
make_tensor("b", &intercept.into_dyn()),
],
doc_string: String::new(),
input: vec![make_value_info(
"X",
&[
Dimension::Symbolic("batch".into()),
Dimension::Fixed(features),
],
)],
output: vec![make_value_info(
"Y",
&[Dimension::Symbolic("batch".into()), Dimension::Fixed(1)],
)],
value_info: Vec::new(),
};
assemble_model(graph, OPSET_VERSION, IR_VERSION)
}