Skip to main content

serde_onnx/ml/
preproc.rs

1use super::{
2    AttrTable, Emitted, OnnxOp, OpError, build_node, check_counts, check_op, push_float,
3    push_floats, push_int, push_ints, push_string, push_strings, push_tensor,
4};
5use crate::ir::{Attribute, ML_DOMAIN, Node, Tensor};
6
7#[derive(Debug, Clone, PartialEq)]
8pub struct Scaler {
9    pub offset: Option<Vec<f32>>,
10    pub scale: Option<Vec<f32>>,
11}
12
13impl OnnxOp for Scaler {
14    const OP_TYPE: &'static str = "Scaler";
15    const DOMAIN: &'static str = ML_DOMAIN;
16
17    fn to_node(&self, inputs: Vec<String>, outputs: Vec<String>) -> Result<Emitted, OpError> {
18        if inputs.len() != 1 {
19            return Err(OpError::WrongInputCount {
20                op: Self::OP_TYPE,
21                expected: "1".to_string(),
22                got: inputs.len(),
23            });
24        }
25        if outputs.len() != 1 {
26            return Err(OpError::WrongOutputCount {
27                op: Self::OP_TYPE,
28                expected: "1".to_string(),
29                got: outputs.len(),
30            });
31        }
32        let mut attrs = Vec::new();
33        push_floats(&mut attrs, "offset", &self.offset);
34        push_floats(&mut attrs, "scale", &self.scale);
35        Ok(build_node::<Self>(attrs, inputs, outputs))
36    }
37
38    fn from_node(node: &Node) -> Result<Self, OpError> {
39        check_op::<Self>(node)?;
40        check_counts(Self::OP_TYPE, node, |n| n == 1, "1", |n| n == 1, "1")?;
41        let mut table = AttrTable::new(Self::OP_TYPE, node)?;
42        let offset = table.opt_floats("offset")?;
43        let scale = table.opt_floats("scale")?;
44        table.finish()?;
45        Ok(Scaler { offset, scale })
46    }
47}
48
49#[derive(Debug, Clone, PartialEq)]
50pub struct Imputer {
51    pub imputed_value_floats: Option<Vec<f32>>,
52    pub imputed_value_int64s: Option<Vec<i64>>,
53    pub replaced_value_float: Option<f32>,
54    pub replaced_value_int64: Option<i64>,
55}
56
57impl OnnxOp for Imputer {
58    const OP_TYPE: &'static str = "Imputer";
59    const DOMAIN: &'static str = ML_DOMAIN;
60
61    fn to_node(&self, inputs: Vec<String>, outputs: Vec<String>) -> Result<Emitted, OpError> {
62        if self.imputed_value_floats.is_some() && self.imputed_value_int64s.is_some() {
63            return Err(OpError::InvalidValue {
64                op: Self::OP_TYPE,
65                attr: "imputed_value_*".to_string(),
66                detail: "only one of imputed_value_floats and imputed_value_int64s may be set"
67                    .to_string(),
68            });
69        }
70        if self.replaced_value_float.is_some() && self.replaced_value_int64.is_some() {
71            return Err(OpError::InvalidValue {
72                op: Self::OP_TYPE,
73                attr: "replaced_value_*".to_string(),
74                detail: "only one of replaced_value_float and replaced_value_int64 may be set"
75                    .to_string(),
76            });
77        }
78        if inputs.len() != 1 {
79            return Err(OpError::WrongInputCount {
80                op: Self::OP_TYPE,
81                expected: "1".to_string(),
82                got: inputs.len(),
83            });
84        }
85        if outputs.len() != 1 {
86            return Err(OpError::WrongOutputCount {
87                op: Self::OP_TYPE,
88                expected: "1".to_string(),
89                got: outputs.len(),
90            });
91        }
92        let mut attrs = Vec::new();
93        push_floats(
94            &mut attrs,
95            "imputed_value_floats",
96            &self.imputed_value_floats,
97        );
98        push_ints(
99            &mut attrs,
100            "imputed_value_int64s",
101            &self.imputed_value_int64s,
102        );
103        push_float(
104            &mut attrs,
105            "replaced_value_float",
106            self.replaced_value_float,
107        );
108        push_int(
109            &mut attrs,
110            "replaced_value_int64",
111            self.replaced_value_int64,
112        );
113        Ok(build_node::<Self>(attrs, inputs, outputs))
114    }
115
116    fn from_node(node: &Node) -> Result<Self, OpError> {
117        check_op::<Self>(node)?;
118        check_counts(Self::OP_TYPE, node, |n| n == 1, "1", |n| n == 1, "1")?;
119        let mut table = AttrTable::new(Self::OP_TYPE, node)?;
120        let imputed_value_floats = table.opt_floats("imputed_value_floats")?;
121        let imputed_value_int64s = table.opt_ints("imputed_value_int64s")?;
122        let replaced_value_float = table.opt_float("replaced_value_float")?;
123        let replaced_value_int64 = table.opt_int("replaced_value_int64")?;
124        table.finish()?;
125        if imputed_value_floats.is_some() && imputed_value_int64s.is_some() {
126            return Err(OpError::InvalidValue {
127                op: Self::OP_TYPE,
128                attr: "imputed_value_*".to_string(),
129                detail: "only one of imputed_value_floats and imputed_value_int64s may be set"
130                    .to_string(),
131            });
132        }
133        if replaced_value_float.is_some() && replaced_value_int64.is_some() {
134            return Err(OpError::InvalidValue {
135                op: Self::OP_TYPE,
136                attr: "replaced_value_*".to_string(),
137                detail: "only one of replaced_value_float and replaced_value_int64 may be set"
138                    .to_string(),
139            });
140        }
141        Ok(Imputer {
142            imputed_value_floats,
143            imputed_value_int64s,
144            replaced_value_float,
145            replaced_value_int64,
146        })
147    }
148}
149
150#[derive(Debug, Clone, PartialEq)]
151pub struct Normalizer {
152    pub norm: Option<String>,
153}
154
155impl OnnxOp for Normalizer {
156    const OP_TYPE: &'static str = "Normalizer";
157    const DOMAIN: &'static str = ML_DOMAIN;
158
159    fn to_node(&self, inputs: Vec<String>, outputs: Vec<String>) -> Result<Emitted, OpError> {
160        if inputs.len() != 1 {
161            return Err(OpError::WrongInputCount {
162                op: Self::OP_TYPE,
163                expected: "1".to_string(),
164                got: inputs.len(),
165            });
166        }
167        if outputs.len() != 1 {
168            return Err(OpError::WrongOutputCount {
169                op: Self::OP_TYPE,
170                expected: "1".to_string(),
171                got: outputs.len(),
172            });
173        }
174        let mut attrs = Vec::new();
175        push_string(&mut attrs, "norm", &self.norm);
176        Ok(build_node::<Self>(attrs, inputs, outputs))
177    }
178
179    fn from_node(node: &Node) -> Result<Self, OpError> {
180        check_op::<Self>(node)?;
181        check_counts(Self::OP_TYPE, node, |n| n == 1, "1", |n| n == 1, "1")?;
182        let mut table = AttrTable::new(Self::OP_TYPE, node)?;
183        let norm = table.opt_string("norm")?;
184        table.finish()?;
185        Ok(Normalizer { norm })
186    }
187}
188
189#[derive(Debug, Clone, PartialEq, Default)]
190pub struct LabelEncoder {
191    pub default_float: Option<f32>,
192    pub default_int64: Option<i64>,
193    pub default_string: Option<String>,
194    pub default_tensor: Option<Tensor>,
195    pub keys_floats: Option<Vec<f32>>,
196    pub keys_int64s: Option<Vec<i64>>,
197    pub keys_strings: Option<Vec<String>>,
198    pub keys_tensor: Option<Tensor>,
199    pub values_floats: Option<Vec<f32>>,
200    pub values_int64s: Option<Vec<i64>>,
201    pub values_strings: Option<Vec<String>>,
202    pub values_tensor: Option<Tensor>,
203}
204
205impl LabelEncoder {
206    fn check_keys_values(&self) -> Result<(), OpError> {
207        let keys = [
208            self.keys_floats.is_some(),
209            self.keys_int64s.is_some(),
210            self.keys_strings.is_some(),
211            self.keys_tensor.is_some(),
212        ]
213        .iter()
214        .filter(|b| **b)
215        .count();
216        if keys != 1 {
217            return Err(OpError::InvalidValue {
218                op: Self::OP_TYPE,
219                attr: "keys_*".to_string(),
220                detail:
221                    "exactly one of keys_floats, keys_int64s, keys_strings, keys_tensor must be set"
222                        .to_string(),
223            });
224        }
225        let values = [
226            self.values_floats.is_some(),
227            self.values_int64s.is_some(),
228            self.values_strings.is_some(),
229            self.values_tensor.is_some(),
230        ]
231        .iter()
232        .filter(|b| **b)
233        .count();
234        if values != 1 {
235            return Err(OpError::InvalidValue {
236                op: Self::OP_TYPE,
237                attr: "values_*".to_string(),
238                detail: "exactly one of values_floats, values_int64s, values_strings, values_tensor must be set"
239                    .to_string(),
240            });
241        }
242        Ok(())
243    }
244}
245
246impl OnnxOp for LabelEncoder {
247    const OP_TYPE: &'static str = "LabelEncoder";
248    const DOMAIN: &'static str = ML_DOMAIN;
249
250    fn to_node(&self, inputs: Vec<String>, outputs: Vec<String>) -> Result<Emitted, OpError> {
251        self.check_keys_values()?;
252        if inputs.len() != 1 {
253            return Err(OpError::WrongInputCount {
254                op: Self::OP_TYPE,
255                expected: "1".to_string(),
256                got: inputs.len(),
257            });
258        }
259        if outputs.len() != 1 {
260            return Err(OpError::WrongOutputCount {
261                op: Self::OP_TYPE,
262                expected: "1".to_string(),
263                got: outputs.len(),
264            });
265        }
266        let mut attrs = Vec::new();
267        push_float(&mut attrs, "default_float", self.default_float);
268        push_int(&mut attrs, "default_int64", self.default_int64);
269        push_string(&mut attrs, "default_string", &self.default_string);
270        push_tensor(&mut attrs, "default_tensor", &self.default_tensor);
271        push_floats(&mut attrs, "keys_floats", &self.keys_floats);
272        push_ints(&mut attrs, "keys_int64s", &self.keys_int64s);
273        push_strings(&mut attrs, "keys_strings", &self.keys_strings);
274        push_tensor(&mut attrs, "keys_tensor", &self.keys_tensor);
275        push_floats(&mut attrs, "values_floats", &self.values_floats);
276        push_ints(&mut attrs, "values_int64s", &self.values_int64s);
277        push_strings(&mut attrs, "values_strings", &self.values_strings);
278        push_tensor(&mut attrs, "values_tensor", &self.values_tensor);
279        Ok(build_node::<Self>(attrs, inputs, outputs))
280    }
281
282    fn from_node(node: &Node) -> Result<Self, OpError> {
283        check_op::<Self>(node)?;
284        check_counts(Self::OP_TYPE, node, |n| n == 1, "1", |n| n == 1, "1")?;
285        let mut table = AttrTable::new(Self::OP_TYPE, node)?;
286        let encoder = LabelEncoder {
287            default_float: table.opt_float("default_float")?,
288            default_int64: table.opt_int("default_int64")?,
289            default_string: table.opt_string("default_string")?,
290            default_tensor: table.opt_tensor("default_tensor")?,
291            keys_floats: table.opt_floats("keys_floats")?,
292            keys_int64s: table.opt_ints("keys_int64s")?,
293            keys_strings: table.opt_strings("keys_strings")?,
294            keys_tensor: table.opt_tensor("keys_tensor")?,
295            values_floats: table.opt_floats("values_floats")?,
296            values_int64s: table.opt_ints("values_int64s")?,
297            values_strings: table.opt_strings("values_strings")?,
298            values_tensor: table.opt_tensor("values_tensor")?,
299        };
300        table.finish()?;
301        encoder.check_keys_values()?;
302        Ok(encoder)
303    }
304}
305
306#[derive(Debug, Clone, PartialEq)]
307pub struct OneHotEncoder {
308    pub cats_int64s: Option<Vec<i64>>,
309    pub cats_strings: Option<Vec<String>>,
310    pub zeros: Option<i64>,
311}
312
313impl OnnxOp for OneHotEncoder {
314    const OP_TYPE: &'static str = "OneHotEncoder";
315    const DOMAIN: &'static str = ML_DOMAIN;
316
317    fn to_node(&self, inputs: Vec<String>, outputs: Vec<String>) -> Result<Emitted, OpError> {
318        check_oneof_cats(
319            Self::OP_TYPE,
320            self.cats_int64s.is_some(),
321            self.cats_strings.is_some(),
322        )?;
323        if inputs.len() != 1 {
324            return Err(OpError::WrongInputCount {
325                op: Self::OP_TYPE,
326                expected: "1".to_string(),
327                got: inputs.len(),
328            });
329        }
330        if outputs.len() != 1 {
331            return Err(OpError::WrongOutputCount {
332                op: Self::OP_TYPE,
333                expected: "1".to_string(),
334                got: outputs.len(),
335            });
336        }
337        let mut attrs = Vec::new();
338        push_ints(&mut attrs, "cats_int64s", &self.cats_int64s);
339        push_strings(&mut attrs, "cats_strings", &self.cats_strings);
340        push_int(&mut attrs, "zeros", self.zeros);
341        Ok(build_node::<Self>(attrs, inputs, outputs))
342    }
343
344    fn from_node(node: &Node) -> Result<Self, OpError> {
345        check_op::<Self>(node)?;
346        check_counts(Self::OP_TYPE, node, |n| n == 1, "1", |n| n == 1, "1")?;
347        let mut table = AttrTable::new(Self::OP_TYPE, node)?;
348        let cats_int64s = table.opt_ints("cats_int64s")?;
349        let cats_strings = table.opt_strings("cats_strings")?;
350        let zeros = table.opt_int("zeros")?;
351        table.finish()?;
352        check_oneof_cats(Self::OP_TYPE, cats_int64s.is_some(), cats_strings.is_some())?;
353        Ok(OneHotEncoder {
354            cats_int64s,
355            cats_strings,
356            zeros,
357        })
358    }
359}
360
361fn check_oneof_cats(op: &'static str, ints: bool, strings: bool) -> Result<(), OpError> {
362    if ints == strings {
363        return Err(OpError::InvalidValue {
364            op,
365            attr: "cats_*".to_string(),
366            detail: "exactly one of cats_int64s and cats_strings must be set".to_string(),
367        });
368    }
369    Ok(())
370}
371
372#[derive(Debug, Clone, PartialEq)]
373pub struct DictVectorizer {
374    pub int64_vocabulary: Option<Vec<i64>>,
375    pub string_vocabulary: Option<Vec<String>>,
376}
377
378impl OnnxOp for DictVectorizer {
379    const OP_TYPE: &'static str = "DictVectorizer";
380    const DOMAIN: &'static str = ML_DOMAIN;
381
382    fn to_node(&self, inputs: Vec<String>, outputs: Vec<String>) -> Result<Emitted, OpError> {
383        check_oneof_vocab(
384            Self::OP_TYPE,
385            self.int64_vocabulary.is_some(),
386            self.string_vocabulary.is_some(),
387        )?;
388        if inputs.len() != 1 {
389            return Err(OpError::WrongInputCount {
390                op: Self::OP_TYPE,
391                expected: "1".to_string(),
392                got: inputs.len(),
393            });
394        }
395        if outputs.len() != 1 {
396            return Err(OpError::WrongOutputCount {
397                op: Self::OP_TYPE,
398                expected: "1".to_string(),
399                got: outputs.len(),
400            });
401        }
402        let mut attrs = Vec::new();
403        push_ints(&mut attrs, "int64_vocabulary", &self.int64_vocabulary);
404        push_strings(&mut attrs, "string_vocabulary", &self.string_vocabulary);
405        Ok(build_node::<Self>(attrs, inputs, outputs))
406    }
407
408    fn from_node(node: &Node) -> Result<Self, OpError> {
409        check_op::<Self>(node)?;
410        check_counts(Self::OP_TYPE, node, |n| n == 1, "1", |n| n == 1, "1")?;
411        let mut table = AttrTable::new(Self::OP_TYPE, node)?;
412        let int64_vocabulary = table.opt_ints("int64_vocabulary")?;
413        let string_vocabulary = table.opt_strings("string_vocabulary")?;
414        table.finish()?;
415        check_oneof_vocab(
416            Self::OP_TYPE,
417            int64_vocabulary.is_some(),
418            string_vocabulary.is_some(),
419        )?;
420        Ok(DictVectorizer {
421            int64_vocabulary,
422            string_vocabulary,
423        })
424    }
425}
426
427fn check_oneof_vocab(op: &'static str, ints: bool, strings: bool) -> Result<(), OpError> {
428    if ints == strings {
429        return Err(OpError::InvalidValue {
430            op,
431            attr: "*_vocabulary".to_string(),
432            detail: "exactly one of int64_vocabulary and string_vocabulary must be set".to_string(),
433        });
434    }
435    Ok(())
436}
437
438#[derive(Debug, Clone, PartialEq)]
439pub struct FeatureVectorizer {
440    pub inputdimensions: Option<Vec<i64>>,
441}
442
443impl OnnxOp for FeatureVectorizer {
444    const OP_TYPE: &'static str = "FeatureVectorizer";
445    const DOMAIN: &'static str = ML_DOMAIN;
446
447    fn to_node(&self, inputs: Vec<String>, outputs: Vec<String>) -> Result<Emitted, OpError> {
448        if inputs.is_empty() {
449            return Err(OpError::WrongInputCount {
450                op: Self::OP_TYPE,
451                expected: "1..".to_string(),
452                got: inputs.len(),
453            });
454        }
455        if outputs.len() != 1 {
456            return Err(OpError::WrongOutputCount {
457                op: Self::OP_TYPE,
458                expected: "1".to_string(),
459                got: outputs.len(),
460            });
461        }
462        let mut attrs: Vec<Attribute> = Vec::new();
463        push_ints(&mut attrs, "inputdimensions", &self.inputdimensions);
464        Ok(build_node::<Self>(attrs, inputs, outputs))
465    }
466
467    fn from_node(node: &Node) -> Result<Self, OpError> {
468        check_op::<Self>(node)?;
469        check_counts(Self::OP_TYPE, node, |n| n >= 1, "1..", |n| n == 1, "1")?;
470        let mut table = AttrTable::new(Self::OP_TYPE, node)?;
471        let inputdimensions = table.opt_ints("inputdimensions")?;
472        table.finish()?;
473        Ok(FeatureVectorizer { inputdimensions })
474    }
475}