Skip to main content

serde_onnx/ml/
mod.rs

1mod aionnx;
2mod linear;
3mod payload;
4mod preproc;
5mod trees;
6
7#[cfg(test)]
8mod tests;
9
10pub use aionnx::{Cast, Concat, Gather, Identity, Reshape};
11pub use linear::{
12    ClassLabels, LinearClassifier, LinearRegressor, SvmClassifier, SvmCommon, SvmRegressor,
13};
14pub use payload::{NodePayload, TypedNode, make_node_payload};
15pub use preproc::{
16    DictVectorizer, FeatureVectorizer, Imputer, LabelEncoder, Normalizer, OneHotEncoder, Scaler,
17};
18pub use trees::{TreeEnsembleClassifier, TreeEnsembleRegressor, TreeNodes, ZipMap};
19
20use crate::ir::{Attribute, AttributeValue, Node, Tensor};
21
22pub const ML_EXPORT_OPSET_TARGET: i64 = 4;
23
24pub const CORE_TYPED_OPSET_TARGET: i64 = 21;
25
26pub const CODEC_MAX_IR_VERSION: i64 = crate::proto::SUPPORTED_IR_VERSION;
27
28pub const CODEC_MAX_ONNX_OPSET: i64 = crate::proto::SUPPORTED_ONNX_OPSET;
29
30pub const CODEC_MAX_ML_OPSET: i64 = crate::proto::SUPPORTED_ML_OPSET;
31
32pub const ML_DOMAIN_STR: &str = crate::ir::ML_DOMAIN;
33
34pub const ONNX_DOMAIN_STR: &str = "";
35
36#[derive(Debug, Clone, PartialEq)]
37pub enum OpError {
38    WrongOp {
39        expected_domain: &'static str,
40        expected_op: &'static str,
41        got_domain: String,
42        got_op: String,
43    },
44    MissingAttribute {
45        op: &'static str,
46        attr: &'static str,
47    },
48    WrongAttributeType {
49        op: &'static str,
50        attr: String,
51        expected: &'static str,
52    },
53    DuplicateAttribute {
54        op: &'static str,
55        attr: String,
56    },
57    UnknownAttribute {
58        op: &'static str,
59        attr: String,
60    },
61    WrongInputCount {
62        op: &'static str,
63        expected: String,
64        got: usize,
65    },
66    WrongOutputCount {
67        op: &'static str,
68        expected: String,
69        got: usize,
70    },
71    InvalidValue {
72        op: &'static str,
73        attr: String,
74        detail: String,
75    },
76}
77
78impl core::fmt::Display for OpError {
79    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
80        match self {
81            Self::WrongOp {
82                expected_domain,
83                expected_op,
84                got_domain,
85                got_op,
86            } => write!(
87                f,
88                "node ({got_domain:?}, {got_op:?}) is not ({expected_domain:?}, {expected_op:?})"
89            ),
90            Self::MissingAttribute { op, attr } => {
91                write!(f, "op {op} is missing required attribute {attr:?}")
92            }
93            Self::WrongAttributeType { op, attr, expected } => write!(
94                f,
95                "op {op} attribute {attr:?} has the wrong type, expected {expected}"
96            ),
97            Self::DuplicateAttribute { op, attr } => {
98                write!(f, "op {op} has duplicate attribute {attr:?}")
99            }
100            Self::UnknownAttribute { op, attr } => {
101                write!(f, "op {op} has unrecognized attribute {attr:?}")
102            }
103            Self::WrongInputCount { op, expected, got } => {
104                write!(f, "op {op} expects {expected} inputs, got {got}")
105            }
106            Self::WrongOutputCount { op, expected, got } => {
107                write!(f, "op {op} expects {expected} outputs, got {got}")
108            }
109            Self::InvalidValue { op, attr, detail } => {
110                write!(f, "op {op} attribute {attr:?} is invalid: {detail}")
111            }
112        }
113    }
114}
115
116impl std::error::Error for OpError {}
117
118#[derive(Debug, Clone, PartialEq)]
119pub struct Emitted {
120    pub node: Node,
121    pub initializers: Vec<Tensor>,
122}
123
124impl Emitted {
125    pub fn single(node: Node) -> Self {
126        Emitted {
127            node,
128            initializers: Vec::new(),
129        }
130    }
131}
132
133pub trait OnnxOp: Sized + Clone + PartialEq + core::fmt::Debug {
134    const OP_TYPE: &'static str;
135    const DOMAIN: &'static str;
136    fn to_node(&self, inputs: Vec<String>, outputs: Vec<String>) -> Result<Emitted, OpError>;
137    fn from_node(node: &Node) -> Result<Self, OpError>;
138}
139
140#[cfg(feature = "export")]
141pub trait OpSink {
142    type Error: core::fmt::Debug;
143    fn emit(&mut self, emitted: Emitted) -> Result<(Vec<String>, Vec<String>), Self::Error>;
144}
145
146#[cfg(feature = "export")]
147pub fn emit_op<S: OpSink, O: OnnxOp>(
148    sink: &mut S,
149    op: &O,
150    inputs: Vec<String>,
151    outputs: Vec<String>,
152) -> Result<(Vec<String>, Vec<String>), S::Error>
153where
154    S::Error: From<OpError>,
155{
156    let emitted = op.to_node(inputs, outputs)?;
157    sink.emit(emitted)
158}
159
160pub(crate) struct AttrTable<'a> {
161    op: &'static str,
162    entries: Vec<(&'a str, &'a AttributeValue)>,
163}
164
165impl<'a> AttrTable<'a> {
166    pub(crate) fn new(op: &'static str, node: &'a Node) -> Result<Self, OpError> {
167        let mut entries = Vec::with_capacity(node.attributes.len());
168        for attr in &node.attributes {
169            if entries.iter().any(|(name, _)| *name == attr.name) {
170                return Err(OpError::DuplicateAttribute {
171                    op,
172                    attr: attr.name.clone(),
173                });
174            }
175            entries.push((attr.name.as_str(), &attr.value));
176        }
177        Ok(AttrTable { op, entries })
178    }
179
180    fn take(&mut self, name: &str) -> Option<&'a AttributeValue> {
181        self.entries
182            .iter()
183            .position(|(n, _)| *n == name)
184            .map(|i| self.entries.remove(i).1)
185    }
186
187    pub(crate) fn opt_floats(&mut self, name: &str) -> Result<Option<Vec<f32>>, OpError> {
188        match self.take(name) {
189            None => Ok(None),
190            Some(AttributeValue::Floats(v)) => Ok(Some(v.clone())),
191            _ => Err(OpError::WrongAttributeType {
192                op: self.op,
193                attr: name.to_string(),
194                expected: "floats",
195            }),
196        }
197    }
198
199    pub(crate) fn opt_ints(&mut self, name: &str) -> Result<Option<Vec<i64>>, OpError> {
200        match self.take(name) {
201            None => Ok(None),
202            Some(AttributeValue::Ints(v)) => Ok(Some(v.clone())),
203            _ => Err(OpError::WrongAttributeType {
204                op: self.op,
205                attr: name.to_string(),
206                expected: "ints",
207            }),
208        }
209    }
210
211    pub(crate) fn opt_strings(&mut self, name: &str) -> Result<Option<Vec<String>>, OpError> {
212        match self.take(name) {
213            None => Ok(None),
214            Some(AttributeValue::Strings(v)) => Ok(Some(v.clone())),
215            _ => Err(OpError::WrongAttributeType {
216                op: self.op,
217                attr: name.to_string(),
218                expected: "strings",
219            }),
220        }
221    }
222
223    pub(crate) fn opt_float(&mut self, name: &str) -> Result<Option<f32>, OpError> {
224        match self.take(name) {
225            None => Ok(None),
226            Some(AttributeValue::Float(v)) => Ok(Some(*v)),
227            _ => Err(OpError::WrongAttributeType {
228                op: self.op,
229                attr: name.to_string(),
230                expected: "float",
231            }),
232        }
233    }
234
235    pub(crate) fn opt_int(&mut self, name: &str) -> Result<Option<i64>, OpError> {
236        match self.take(name) {
237            None => Ok(None),
238            Some(AttributeValue::Int(v)) => Ok(Some(*v)),
239            _ => Err(OpError::WrongAttributeType {
240                op: self.op,
241                attr: name.to_string(),
242                expected: "int",
243            }),
244        }
245    }
246
247    pub(crate) fn opt_string(&mut self, name: &str) -> Result<Option<String>, OpError> {
248        match self.take(name) {
249            None => Ok(None),
250            Some(AttributeValue::String(v)) => Ok(Some(v.clone())),
251            _ => Err(OpError::WrongAttributeType {
252                op: self.op,
253                attr: name.to_string(),
254                expected: "string",
255            }),
256        }
257    }
258
259    pub(crate) fn opt_tensor(&mut self, name: &str) -> Result<Option<Tensor>, OpError> {
260        match self.take(name) {
261            None => Ok(None),
262            Some(AttributeValue::Tensor(v)) => Ok(Some((**v).clone())),
263            _ => Err(OpError::WrongAttributeType {
264                op: self.op,
265                attr: name.to_string(),
266                expected: "tensor",
267            }),
268        }
269    }
270
271    pub(crate) fn req_floats(&mut self, name: &'static str) -> Result<Vec<f32>, OpError> {
272        self.opt_floats(name)?.ok_or(OpError::MissingAttribute {
273            op: self.op,
274            attr: name,
275        })
276    }
277
278    pub(crate) fn req_int(&mut self, name: &'static str) -> Result<i64, OpError> {
279        self.opt_int(name)?.ok_or(OpError::MissingAttribute {
280            op: self.op,
281            attr: name,
282        })
283    }
284
285    pub(crate) fn finish(self) -> Result<(), OpError> {
286        match self.entries.into_iter().next() {
287            None => Ok(()),
288            Some((name, _)) => Err(OpError::UnknownAttribute {
289                op: self.op,
290                attr: name.to_string(),
291            }),
292        }
293    }
294}
295
296pub(crate) fn check_op<O: OnnxOp>(node: &Node) -> Result<(), OpError> {
297    if node.domain != O::DOMAIN || node.op_type != O::OP_TYPE {
298        return Err(OpError::WrongOp {
299            expected_domain: O::DOMAIN,
300            expected_op: O::OP_TYPE,
301            got_domain: node.domain.clone(),
302            got_op: node.op_type.clone(),
303        });
304    }
305    Ok(())
306}
307
308pub(crate) fn check_counts(
309    op: &'static str,
310    node: &Node,
311    inputs: fn(usize) -> bool,
312    inputs_desc: &str,
313    outputs: fn(usize) -> bool,
314    outputs_desc: &str,
315) -> Result<(), OpError> {
316    if !inputs(node.inputs.len()) {
317        return Err(OpError::WrongInputCount {
318            op,
319            expected: inputs_desc.to_string(),
320            got: node.inputs.len(),
321        });
322    }
323    if !outputs(node.outputs.len()) {
324        return Err(OpError::WrongOutputCount {
325            op,
326            expected: outputs_desc.to_string(),
327            got: node.outputs.len(),
328        });
329    }
330    Ok(())
331}
332
333pub(crate) fn build_node<O: OnnxOp>(
334    op_attrs: Vec<Attribute>,
335    inputs: Vec<String>,
336    outputs: Vec<String>,
337) -> Emitted {
338    Emitted::single(Node::new(O::OP_TYPE, O::DOMAIN, inputs, outputs, op_attrs))
339}
340
341pub(crate) fn push_floats(attrs: &mut Vec<Attribute>, name: &str, v: &Option<Vec<f32>>) {
342    if let Some(x) = v {
343        attrs.push(Attribute::floats(name, x.clone()));
344    }
345}
346
347pub(crate) fn push_ints(attrs: &mut Vec<Attribute>, name: &str, v: &Option<Vec<i64>>) {
348    if let Some(x) = v {
349        attrs.push(Attribute::ints(name, x.clone()));
350    }
351}
352
353pub(crate) fn push_strings(attrs: &mut Vec<Attribute>, name: &str, v: &Option<Vec<String>>) {
354    if let Some(x) = v {
355        attrs.push(Attribute::strings(name, x.clone()));
356    }
357}
358
359pub(crate) fn push_float(attrs: &mut Vec<Attribute>, name: &str, v: Option<f32>) {
360    if let Some(x) = v {
361        attrs.push(Attribute::float(name, x));
362    }
363}
364
365pub(crate) fn push_int(attrs: &mut Vec<Attribute>, name: &str, v: Option<i64>) {
366    if let Some(x) = v {
367        attrs.push(Attribute::int(name, x));
368    }
369}
370
371pub(crate) fn push_string(attrs: &mut Vec<Attribute>, name: &str, v: &Option<String>) {
372    if let Some(x) = v {
373        attrs.push(Attribute::string(name, x.clone()));
374    }
375}
376
377pub(crate) fn push_tensor(attrs: &mut Vec<Attribute>, name: &str, v: &Option<Tensor>) {
378    if let Some(x) = v {
379        attrs.push(Attribute::tensor(name, x.clone()));
380    }
381}