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}