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}