1use super::{
2 AttrTable, Emitted, OnnxOp, OpError, build_node, check_counts, check_op, push_floats, push_int,
3 push_ints, push_string, push_strings, push_tensor,
4};
5use crate::ir::{Attribute, ML_DOMAIN, Node, Tensor};
6
7#[derive(Debug, Clone, PartialEq, Default)]
8pub struct TreeNodes {
9 pub falsenodeids: Vec<i64>,
10 pub featureids: Vec<i64>,
11 pub hitrates: Option<Vec<f32>>,
12 pub hitrates_as_tensor: Option<Tensor>,
13 pub missing_value_tracks_true: Option<Vec<i64>>,
14 pub modes: Vec<String>,
15 pub nodeids: Vec<i64>,
16 pub treeids: Vec<i64>,
17 pub truenodeids: Vec<i64>,
18 pub values: Option<Vec<f32>>,
19 pub values_as_tensor: Option<Tensor>,
20}
21
22impl TreeNodes {
23 fn check(&self, op: &'static str) -> Result<(), OpError> {
24 let n = self.modes.len();
25 if n == 0 {
26 return Err(OpError::MissingAttribute {
27 op,
28 attr: "nodes_modes",
29 });
30 }
31 for (attr, len) in [
32 ("nodes_falsenodeids", self.falsenodeids.len()),
33 ("nodes_featureids", self.featureids.len()),
34 ("nodes_nodeids", self.nodeids.len()),
35 ("nodes_treeids", self.treeids.len()),
36 ("nodes_truenodeids", self.truenodeids.len()),
37 ] {
38 if len != n {
39 return Err(OpError::InvalidValue {
40 op,
41 attr: attr.to_string(),
42 detail: "all nodes_* arrays must share the length of nodes_modes".to_string(),
43 });
44 }
45 }
46 if let Some(v) = &self.missing_value_tracks_true
47 && v.len() != n
48 {
49 return Err(OpError::InvalidValue {
50 op,
51 attr: "nodes_missing_value_tracks_true".to_string(),
52 detail: "all nodes_* arrays must share the length of nodes_modes".to_string(),
53 });
54 }
55 if let Some(v) = &self.hitrates
56 && v.len() != n
57 {
58 return Err(OpError::InvalidValue {
59 op,
60 attr: "nodes_hitrates".to_string(),
61 detail: "all nodes_* arrays must share the length of nodes_modes".to_string(),
62 });
63 }
64 if let Some(v) = &self.values
65 && v.len() != n
66 {
67 return Err(OpError::InvalidValue {
68 op,
69 attr: "nodes_values".to_string(),
70 detail: "all nodes_* arrays must share the length of nodes_modes".to_string(),
71 });
72 }
73 Ok(())
74 }
75
76 fn emit(&self, attrs: &mut Vec<Attribute>) {
77 push_ints(
78 attrs,
79 "nodes_falsenodeids",
80 &Some(self.falsenodeids.clone()),
81 );
82 push_ints(attrs, "nodes_featureids", &Some(self.featureids.clone()));
83 push_floats(attrs, "nodes_hitrates", &self.hitrates);
84 push_tensor(attrs, "nodes_hitrates_as_tensor", &self.hitrates_as_tensor);
85 push_ints(
86 attrs,
87 "nodes_missing_value_tracks_true",
88 &self.missing_value_tracks_true,
89 );
90 push_strings(attrs, "nodes_modes", &Some(self.modes.clone()));
91 push_ints(attrs, "nodes_nodeids", &Some(self.nodeids.clone()));
92 push_ints(attrs, "nodes_treeids", &Some(self.treeids.clone()));
93 push_ints(attrs, "nodes_truenodeids", &Some(self.truenodeids.clone()));
94 push_floats(attrs, "nodes_values", &self.values);
95 push_tensor(attrs, "nodes_values_as_tensor", &self.values_as_tensor);
96 }
97
98 fn parse(op: &'static str, table: &mut AttrTable) -> Result<Self, OpError> {
99 let nodes = TreeNodes {
100 falsenodeids: table.opt_ints("nodes_falsenodeids")?.unwrap_or_default(),
101 featureids: table.opt_ints("nodes_featureids")?.unwrap_or_default(),
102 hitrates: table.opt_floats("nodes_hitrates")?,
103 hitrates_as_tensor: table.opt_tensor("nodes_hitrates_as_tensor")?,
104 missing_value_tracks_true: table.opt_ints("nodes_missing_value_tracks_true")?,
105 modes: table.opt_strings("nodes_modes")?.unwrap_or_default(),
106 nodeids: table.opt_ints("nodes_nodeids")?.unwrap_or_default(),
107 treeids: table.opt_ints("nodes_treeids")?.unwrap_or_default(),
108 truenodeids: table.opt_ints("nodes_truenodeids")?.unwrap_or_default(),
109 values: table.opt_floats("nodes_values")?,
110 values_as_tensor: table.opt_tensor("nodes_values_as_tensor")?,
111 };
112 nodes.check(op)?;
113 Ok(nodes)
114 }
115}
116
117#[derive(Debug, Clone, PartialEq)]
118pub struct TreeEnsembleClassifier {
119 pub base_values: Option<Vec<f32>>,
120 pub base_values_as_tensor: Option<Tensor>,
121 pub class_ids: Vec<i64>,
122 pub class_nodeids: Vec<i64>,
123 pub class_treeids: Vec<i64>,
124 pub class_weights: Option<Vec<f32>>,
125 pub class_weights_as_tensor: Option<Tensor>,
126 pub classlabels_int64s: Option<Vec<i64>>,
127 pub classlabels_strings: Option<Vec<String>>,
128 pub nodes: TreeNodes,
129 pub post_transform: Option<String>,
130}
131
132impl TreeEnsembleClassifier {
133 fn check(&self) -> Result<(), OpError> {
134 const OP: &str = "TreeEnsembleClassifier";
135 if self.classlabels_int64s.is_some() == self.classlabels_strings.is_some() {
136 return Err(OpError::InvalidValue {
137 op: OP,
138 attr: "classlabels_*".to_string(),
139 detail: "exactly one of classlabels_int64s and classlabels_strings must be set"
140 .to_string(),
141 });
142 }
143 self.nodes.check(OP)?;
144 let n = self.class_ids.len();
145 if n == 0 {
146 return Err(OpError::MissingAttribute {
147 op: OP,
148 attr: "class_ids",
149 });
150 }
151 for (attr, len) in [
152 ("class_nodeids", self.class_nodeids.len()),
153 ("class_treeids", self.class_treeids.len()),
154 ] {
155 if len != n {
156 return Err(OpError::InvalidValue {
157 op: OP,
158 attr: attr.to_string(),
159 detail: "all class_* arrays must share the length of class_ids".to_string(),
160 });
161 }
162 }
163 if let Some(v) = &self.class_weights
164 && v.len() != n
165 {
166 return Err(OpError::InvalidValue {
167 op: OP,
168 attr: "class_weights".to_string(),
169 detail: "all class_* arrays must share the length of class_ids".to_string(),
170 });
171 }
172 Ok(())
173 }
174}
175
176impl OnnxOp for TreeEnsembleClassifier {
177 const OP_TYPE: &'static str = "TreeEnsembleClassifier";
178 const DOMAIN: &'static str = ML_DOMAIN;
179
180 fn to_node(&self, inputs: Vec<String>, outputs: Vec<String>) -> Result<Emitted, OpError> {
181 self.check()?;
182 if inputs.len() != 1 {
183 return Err(OpError::WrongInputCount {
184 op: Self::OP_TYPE,
185 expected: "1".to_string(),
186 got: inputs.len(),
187 });
188 }
189 if outputs.len() != 2 {
190 return Err(OpError::WrongOutputCount {
191 op: Self::OP_TYPE,
192 expected: "2".to_string(),
193 got: outputs.len(),
194 });
195 }
196 let mut attrs = Vec::new();
197 push_floats(&mut attrs, "base_values", &self.base_values);
198 push_tensor(
199 &mut attrs,
200 "base_values_as_tensor",
201 &self.base_values_as_tensor,
202 );
203 push_ints(&mut attrs, "class_ids", &Some(self.class_ids.clone()));
204 push_ints(
205 &mut attrs,
206 "class_nodeids",
207 &Some(self.class_nodeids.clone()),
208 );
209 push_ints(
210 &mut attrs,
211 "class_treeids",
212 &Some(self.class_treeids.clone()),
213 );
214 push_floats(&mut attrs, "class_weights", &self.class_weights);
215 push_tensor(
216 &mut attrs,
217 "class_weights_as_tensor",
218 &self.class_weights_as_tensor,
219 );
220 push_ints(&mut attrs, "classlabels_int64s", &self.classlabels_int64s);
221 push_strings(&mut attrs, "classlabels_strings", &self.classlabels_strings);
222 self.nodes.emit(&mut attrs);
223 push_string(&mut attrs, "post_transform", &self.post_transform);
224 Ok(build_node::<Self>(attrs, inputs, outputs))
225 }
226
227 fn from_node(node: &Node) -> Result<Self, OpError> {
228 check_op::<Self>(node)?;
229 check_counts(Self::OP_TYPE, node, |n| n == 1, "1", |n| n == 2, "2")?;
230 let mut table = AttrTable::new(Self::OP_TYPE, node)?;
231 let ensemble = TreeEnsembleClassifier {
232 base_values: table.opt_floats("base_values")?,
233 base_values_as_tensor: table.opt_tensor("base_values_as_tensor")?,
234 class_ids: table.opt_ints("class_ids")?.unwrap_or_default(),
235 class_nodeids: table.opt_ints("class_nodeids")?.unwrap_or_default(),
236 class_treeids: table.opt_ints("class_treeids")?.unwrap_or_default(),
237 class_weights: table.opt_floats("class_weights")?,
238 class_weights_as_tensor: table.opt_tensor("class_weights_as_tensor")?,
239 classlabels_int64s: table.opt_ints("classlabels_int64s")?,
240 classlabels_strings: table.opt_strings("classlabels_strings")?,
241 nodes: TreeNodes::parse(Self::OP_TYPE, &mut table)?,
242 post_transform: table.opt_string("post_transform")?,
243 };
244 table.finish()?;
245 ensemble.check()?;
246 Ok(ensemble)
247 }
248}
249
250#[derive(Debug, Clone, PartialEq)]
251pub struct TreeEnsembleRegressor {
252 pub aggregate_function: Option<String>,
253 pub base_values: Option<Vec<f32>>,
254 pub base_values_as_tensor: Option<Tensor>,
255 pub n_targets: Option<i64>,
256 pub nodes: TreeNodes,
257 pub post_transform: Option<String>,
258 pub target_ids: Vec<i64>,
259 pub target_nodeids: Vec<i64>,
260 pub target_treeids: Vec<i64>,
261 pub target_weights: Option<Vec<f32>>,
262 pub target_weights_as_tensor: Option<Tensor>,
263}
264
265impl TreeEnsembleRegressor {
266 fn check(&self) -> Result<(), OpError> {
267 const OP: &str = "TreeEnsembleRegressor";
268 self.nodes.check(OP)?;
269 let n = self.target_ids.len();
270 if n == 0 {
271 return Err(OpError::MissingAttribute {
272 op: OP,
273 attr: "target_ids",
274 });
275 }
276 for (attr, len) in [
277 ("target_nodeids", self.target_nodeids.len()),
278 ("target_treeids", self.target_treeids.len()),
279 ] {
280 if len != n {
281 return Err(OpError::InvalidValue {
282 op: OP,
283 attr: attr.to_string(),
284 detail: "all target_* arrays must share the length of target_ids".to_string(),
285 });
286 }
287 }
288 if let Some(v) = &self.target_weights
289 && v.len() != n
290 {
291 return Err(OpError::InvalidValue {
292 op: OP,
293 attr: "target_weights".to_string(),
294 detail: "all target_* arrays must share the length of target_ids".to_string(),
295 });
296 }
297 Ok(())
298 }
299}
300
301impl OnnxOp for TreeEnsembleRegressor {
302 const OP_TYPE: &'static str = "TreeEnsembleRegressor";
303 const DOMAIN: &'static str = ML_DOMAIN;
304
305 fn to_node(&self, inputs: Vec<String>, outputs: Vec<String>) -> Result<Emitted, OpError> {
306 self.check()?;
307 if inputs.len() != 1 {
308 return Err(OpError::WrongInputCount {
309 op: Self::OP_TYPE,
310 expected: "1".to_string(),
311 got: inputs.len(),
312 });
313 }
314 if outputs.len() != 1 {
315 return Err(OpError::WrongOutputCount {
316 op: Self::OP_TYPE,
317 expected: "1".to_string(),
318 got: outputs.len(),
319 });
320 }
321 let mut attrs = Vec::new();
322 push_string(&mut attrs, "aggregate_function", &self.aggregate_function);
323 push_floats(&mut attrs, "base_values", &self.base_values);
324 push_tensor(
325 &mut attrs,
326 "base_values_as_tensor",
327 &self.base_values_as_tensor,
328 );
329 push_int(&mut attrs, "n_targets", self.n_targets);
330 self.nodes.emit(&mut attrs);
331 push_string(&mut attrs, "post_transform", &self.post_transform);
332 push_ints(&mut attrs, "target_ids", &Some(self.target_ids.clone()));
333 push_ints(
334 &mut attrs,
335 "target_nodeids",
336 &Some(self.target_nodeids.clone()),
337 );
338 push_ints(
339 &mut attrs,
340 "target_treeids",
341 &Some(self.target_treeids.clone()),
342 );
343 push_floats(&mut attrs, "target_weights", &self.target_weights);
344 push_tensor(
345 &mut attrs,
346 "target_weights_as_tensor",
347 &self.target_weights_as_tensor,
348 );
349 Ok(build_node::<Self>(attrs, inputs, outputs))
350 }
351
352 fn from_node(node: &Node) -> Result<Self, OpError> {
353 check_op::<Self>(node)?;
354 check_counts(Self::OP_TYPE, node, |n| n == 1, "1", |n| n == 1, "1")?;
355 let mut table = AttrTable::new(Self::OP_TYPE, node)?;
356 let ensemble = TreeEnsembleRegressor {
357 aggregate_function: table.opt_string("aggregate_function")?,
358 base_values: table.opt_floats("base_values")?,
359 base_values_as_tensor: table.opt_tensor("base_values_as_tensor")?,
360 n_targets: table.opt_int("n_targets")?,
361 nodes: TreeNodes::parse(Self::OP_TYPE, &mut table)?,
362 post_transform: table.opt_string("post_transform")?,
363 target_ids: table.opt_ints("target_ids")?.unwrap_or_default(),
364 target_nodeids: table.opt_ints("target_nodeids")?.unwrap_or_default(),
365 target_treeids: table.opt_ints("target_treeids")?.unwrap_or_default(),
366 target_weights: table.opt_floats("target_weights")?,
367 target_weights_as_tensor: table.opt_tensor("target_weights_as_tensor")?,
368 };
369 table.finish()?;
370 ensemble.check()?;
371 Ok(ensemble)
372 }
373}
374
375#[derive(Debug, Clone, PartialEq)]
376pub struct ZipMap {
377 pub classlabels_int64s: Option<Vec<i64>>,
378 pub classlabels_strings: Option<Vec<String>>,
379}
380
381impl OnnxOp for ZipMap {
382 const OP_TYPE: &'static str = "ZipMap";
383 const DOMAIN: &'static str = ML_DOMAIN;
384
385 fn to_node(&self, inputs: Vec<String>, outputs: Vec<String>) -> Result<Emitted, OpError> {
386 if self.classlabels_int64s.is_some() == self.classlabels_strings.is_some() {
387 return Err(OpError::InvalidValue {
388 op: Self::OP_TYPE,
389 attr: "classlabels_*".to_string(),
390 detail: "exactly one of classlabels_int64s and classlabels_strings must be set"
391 .to_string(),
392 });
393 }
394 if inputs.len() != 1 {
395 return Err(OpError::WrongInputCount {
396 op: Self::OP_TYPE,
397 expected: "1".to_string(),
398 got: inputs.len(),
399 });
400 }
401 if outputs.len() != 1 {
402 return Err(OpError::WrongOutputCount {
403 op: Self::OP_TYPE,
404 expected: "1".to_string(),
405 got: outputs.len(),
406 });
407 }
408 let mut attrs = Vec::new();
409 push_ints(&mut attrs, "classlabels_int64s", &self.classlabels_int64s);
410 push_strings(&mut attrs, "classlabels_strings", &self.classlabels_strings);
411 Ok(build_node::<Self>(attrs, inputs, outputs))
412 }
413
414 fn from_node(node: &Node) -> Result<Self, OpError> {
415 check_op::<Self>(node)?;
416 check_counts(Self::OP_TYPE, node, |n| n == 1, "1", |n| n == 1, "1")?;
417 let mut table = AttrTable::new(Self::OP_TYPE, node)?;
418 let map = ZipMap {
419 classlabels_int64s: table.opt_ints("classlabels_int64s")?,
420 classlabels_strings: table.opt_strings("classlabels_strings")?,
421 };
422 table.finish()?;
423 if map.classlabels_int64s.is_some() == map.classlabels_strings.is_some() {
424 return Err(OpError::InvalidValue {
425 op: Self::OP_TYPE,
426 attr: "classlabels_*".to_string(),
427 detail: "exactly one of classlabels_int64s and classlabels_strings must be set"
428 .to_string(),
429 });
430 }
431 Ok(map)
432 }
433}