use ndarray::{Array1, Array2};
use crate::canonical::LinearModelWeights;
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_linear(weights: &LinearModelWeights) -> ModelProto {
let features = weights.n_features();
let coefficients = Array2::from_shape_fn((features, 1), |(feature, _)| {
weights.coefficients[feature] as f32
});
let intercept = Array1::from_vec(vec![weights.intercept as f32]);
let graph = GraphProto {
node: vec![make_node("Gemm", ["X", "W", "b"], ["Y"], Vec::new())],
name: "linear_regression".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)
}