Skip to main content

serde_onnx/ml/
trees.rs

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}