Skip to main content

serde_onnx/ml/
payload.rs

1use super::{
2    Cast, Concat, DictVectorizer, FeatureVectorizer, Gather, Identity, Imputer, LabelEncoder,
3    LinearClassifier, LinearRegressor, Normalizer, OneHotEncoder, OnnxOp, Reshape, Scaler,
4    SvmClassifier, SvmRegressor, TreeEnsembleClassifier, TreeEnsembleRegressor, ZipMap,
5};
6use crate::ir::Node;
7
8#[derive(Debug, Clone, PartialEq)]
9pub struct TypedNode<O> {
10    pub op: O,
11    pub node: Node,
12}
13
14impl<O: OnnxOp> TypedNode<O> {
15    pub fn recognize(node: &Node) -> Result<Self, super::OpError> {
16        Ok(TypedNode {
17            op: O::from_node(node)?,
18            node: node.clone(),
19        })
20    }
21}
22
23#[derive(Debug, Clone, PartialEq)]
24pub enum NodePayload {
25    Scaler(TypedNode<Scaler>),
26    Imputer(TypedNode<Imputer>),
27    Normalizer(TypedNode<Normalizer>),
28    LabelEncoder(TypedNode<LabelEncoder>),
29    OneHotEncoder(TypedNode<OneHotEncoder>),
30    DictVectorizer(TypedNode<DictVectorizer>),
31    FeatureVectorizer(TypedNode<FeatureVectorizer>),
32    LinearClassifier(TypedNode<LinearClassifier>),
33    LinearRegressor(TypedNode<LinearRegressor>),
34    SvmClassifier(TypedNode<SvmClassifier>),
35    SvmRegressor(TypedNode<SvmRegressor>),
36    TreeEnsembleClassifier(TypedNode<TreeEnsembleClassifier>),
37    TreeEnsembleRegressor(TypedNode<TreeEnsembleRegressor>),
38    ZipMap(TypedNode<ZipMap>),
39    Cast(TypedNode<Cast>),
40    Reshape(TypedNode<Reshape>),
41    Concat(TypedNode<Concat>),
42    Gather(TypedNode<Gather>),
43    Identity(TypedNode<Identity>),
44    Raw(Node),
45}
46
47impl NodePayload {
48    pub fn node(&self) -> &Node {
49        match self {
50            Self::Scaler(v) => &v.node,
51            Self::Imputer(v) => &v.node,
52            Self::Normalizer(v) => &v.node,
53            Self::LabelEncoder(v) => &v.node,
54            Self::OneHotEncoder(v) => &v.node,
55            Self::DictVectorizer(v) => &v.node,
56            Self::FeatureVectorizer(v) => &v.node,
57            Self::LinearClassifier(v) => &v.node,
58            Self::LinearRegressor(v) => &v.node,
59            Self::SvmClassifier(v) => &v.node,
60            Self::SvmRegressor(v) => &v.node,
61            Self::TreeEnsembleClassifier(v) => &v.node,
62            Self::TreeEnsembleRegressor(v) => &v.node,
63            Self::ZipMap(v) => &v.node,
64            Self::Cast(v) => &v.node,
65            Self::Reshape(v) => &v.node,
66            Self::Concat(v) => &v.node,
67            Self::Gather(v) => &v.node,
68            Self::Identity(v) => &v.node,
69            Self::Raw(node) => node,
70        }
71    }
72
73    pub fn is_raw(&self) -> bool {
74        matches!(self, Self::Raw(_))
75    }
76}
77
78pub fn make_node_payload(node: &Node) -> NodePayload {
79    const ML: &str = crate::ir::ML_DOMAIN;
80    const CORE: &str = super::ONNX_DOMAIN_STR;
81    match (node.domain.as_str(), node.op_type.as_str()) {
82        (ML, "Scaler") => TypedNode::recognize(node).map(NodePayload::Scaler),
83        (ML, "Imputer") => TypedNode::recognize(node).map(NodePayload::Imputer),
84        (ML, "Normalizer") => TypedNode::recognize(node).map(NodePayload::Normalizer),
85        (ML, "LabelEncoder") => TypedNode::recognize(node).map(NodePayload::LabelEncoder),
86        (ML, "OneHotEncoder") => TypedNode::recognize(node).map(NodePayload::OneHotEncoder),
87        (ML, "DictVectorizer") => TypedNode::recognize(node).map(NodePayload::DictVectorizer),
88        (ML, "FeatureVectorizer") => TypedNode::recognize(node).map(NodePayload::FeatureVectorizer),
89        (ML, "LinearClassifier") => TypedNode::recognize(node).map(NodePayload::LinearClassifier),
90        (ML, "LinearRegressor") => TypedNode::recognize(node).map(NodePayload::LinearRegressor),
91        (ML, "SVMClassifier") => TypedNode::recognize(node).map(NodePayload::SvmClassifier),
92        (ML, "SVMRegressor") => TypedNode::recognize(node).map(NodePayload::SvmRegressor),
93        (ML, "TreeEnsembleClassifier") => {
94            TypedNode::recognize(node).map(NodePayload::TreeEnsembleClassifier)
95        }
96        (ML, "TreeEnsembleRegressor") => {
97            TypedNode::recognize(node).map(NodePayload::TreeEnsembleRegressor)
98        }
99        (ML, "ZipMap") => TypedNode::recognize(node).map(NodePayload::ZipMap),
100        (CORE, "Cast") => TypedNode::recognize(node).map(NodePayload::Cast),
101        (CORE, "Reshape") => TypedNode::recognize(node).map(NodePayload::Reshape),
102        (CORE, "Concat") => TypedNode::recognize(node).map(NodePayload::Concat),
103        (CORE, "Gather") => TypedNode::recognize(node).map(NodePayload::Gather),
104        (CORE, "Identity") => TypedNode::recognize(node).map(NodePayload::Identity),
105        _ => Err(super::OpError::WrongOp {
106            expected_domain: "known",
107            expected_op: "known",
108            got_domain: node.domain.clone(),
109            got_op: node.op_type.clone(),
110        }),
111    }
112    .unwrap_or_else(|_| NodePayload::Raw(node.clone()))
113}