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}