Skip to main content

serde_onnx/ml/
linear.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,
4};
5use crate::ir::{ML_DOMAIN, Node};
6
7#[derive(Debug, Clone, PartialEq, Default)]
8pub struct ClassLabels {
9    pub ints: Option<Vec<i64>>,
10    pub strings: Option<Vec<String>>,
11}
12
13impl ClassLabels {
14    fn check(&self, op: &'static str) -> Result<(), OpError> {
15        if self.ints.is_some() == self.strings.is_some() {
16            return Err(OpError::InvalidValue {
17                op,
18                attr: "classlabels_*".to_string(),
19                detail: "exactly one of classlabels_ints and classlabels_strings must be set"
20                    .to_string(),
21            });
22        }
23        Ok(())
24    }
25
26    fn emit(&self, attrs: &mut Vec<crate::ir::Attribute>) {
27        push_ints(attrs, "classlabels_ints", &self.ints);
28        push_strings(attrs, "classlabels_strings", &self.strings);
29    }
30
31    fn parse(op: &'static str, table: &mut AttrTable) -> Result<Self, OpError> {
32        let labels = ClassLabels {
33            ints: table.opt_ints("classlabels_ints")?,
34            strings: table.opt_strings("classlabels_strings")?,
35        };
36        labels.check(op)?;
37        Ok(labels)
38    }
39}
40
41#[derive(Debug, Clone, PartialEq)]
42pub struct LinearClassifier {
43    pub classlabels: ClassLabels,
44    pub coefficients: Vec<f32>,
45    pub intercepts: Option<Vec<f32>>,
46    pub multi_class: Option<i64>,
47    pub post_transform: Option<String>,
48}
49
50impl LinearClassifier {
51    pub fn with_int_labels(labels: Vec<i64>, coefficients: Vec<f32>) -> Result<Self, OpError> {
52        let classifier = LinearClassifier {
53            classlabels: ClassLabels {
54                ints: Some(labels),
55                strings: None,
56            },
57            coefficients,
58            intercepts: None,
59            multi_class: None,
60            post_transform: None,
61        };
62        classifier.check()?;
63        Ok(classifier)
64    }
65
66    pub fn with_string_labels(
67        labels: Vec<String>,
68        coefficients: Vec<f32>,
69    ) -> Result<Self, OpError> {
70        let classifier = LinearClassifier {
71            classlabels: ClassLabels {
72                ints: None,
73                strings: Some(labels),
74            },
75            coefficients,
76            intercepts: None,
77            multi_class: None,
78            post_transform: None,
79        };
80        classifier.check()?;
81        Ok(classifier)
82    }
83
84    fn check(&self) -> Result<(), OpError> {
85        self.classlabels.check(Self::OP_TYPE)?;
86        if self.coefficients.is_empty() {
87            return Err(OpError::MissingAttribute {
88                op: Self::OP_TYPE,
89                attr: "coefficients",
90            });
91        }
92        Ok(())
93    }
94}
95
96impl OnnxOp for LinearClassifier {
97    const OP_TYPE: &'static str = "LinearClassifier";
98    const DOMAIN: &'static str = ML_DOMAIN;
99
100    fn to_node(&self, inputs: Vec<String>, outputs: Vec<String>) -> Result<Emitted, OpError> {
101        self.check()?;
102        if inputs.len() != 1 {
103            return Err(OpError::WrongInputCount {
104                op: Self::OP_TYPE,
105                expected: "1".to_string(),
106                got: inputs.len(),
107            });
108        }
109        if outputs.len() != 2 {
110            return Err(OpError::WrongOutputCount {
111                op: Self::OP_TYPE,
112                expected: "2".to_string(),
113                got: outputs.len(),
114            });
115        }
116        let mut attrs = Vec::new();
117        self.classlabels.emit(&mut attrs);
118        attrs.push(crate::ir::Attribute::floats(
119            "coefficients",
120            self.coefficients.clone(),
121        ));
122        push_floats(&mut attrs, "intercepts", &self.intercepts);
123        push_int(&mut attrs, "multi_class", self.multi_class);
124        push_string(&mut attrs, "post_transform", &self.post_transform);
125        Ok(build_node::<Self>(attrs, inputs, outputs))
126    }
127
128    fn from_node(node: &Node) -> Result<Self, OpError> {
129        check_op::<Self>(node)?;
130        check_counts(Self::OP_TYPE, node, |n| n == 1, "1", |n| n == 2, "2")?;
131        let mut table = AttrTable::new(Self::OP_TYPE, node)?;
132        let classlabels = ClassLabels::parse(Self::OP_TYPE, &mut table)?;
133        let coefficients = table.req_floats("coefficients")?;
134        let intercepts = table.opt_floats("intercepts")?;
135        let multi_class = table.opt_int("multi_class")?;
136        let post_transform = table.opt_string("post_transform")?;
137        table.finish()?;
138        let classifier = LinearClassifier {
139            classlabels,
140            coefficients,
141            intercepts,
142            multi_class,
143            post_transform,
144        };
145        classifier.check()?;
146        Ok(classifier)
147    }
148}
149
150#[derive(Debug, Clone, PartialEq)]
151pub struct LinearRegressor {
152    pub coefficients: Vec<f32>,
153    pub intercepts: Option<Vec<f32>>,
154    pub post_transform: Option<String>,
155    pub targets: Option<i64>,
156}
157
158impl OnnxOp for LinearRegressor {
159    const OP_TYPE: &'static str = "LinearRegressor";
160    const DOMAIN: &'static str = ML_DOMAIN;
161
162    fn to_node(&self, inputs: Vec<String>, outputs: Vec<String>) -> Result<Emitted, OpError> {
163        if self.coefficients.is_empty() {
164            return Err(OpError::MissingAttribute {
165                op: Self::OP_TYPE,
166                attr: "coefficients",
167            });
168        }
169        if inputs.len() != 1 {
170            return Err(OpError::WrongInputCount {
171                op: Self::OP_TYPE,
172                expected: "1".to_string(),
173                got: inputs.len(),
174            });
175        }
176        if outputs.len() != 1 {
177            return Err(OpError::WrongOutputCount {
178                op: Self::OP_TYPE,
179                expected: "1".to_string(),
180                got: outputs.len(),
181            });
182        }
183        let mut attrs = Vec::new();
184        attrs.push(crate::ir::Attribute::floats(
185            "coefficients",
186            self.coefficients.clone(),
187        ));
188        push_floats(&mut attrs, "intercepts", &self.intercepts);
189        push_string(&mut attrs, "post_transform", &self.post_transform);
190        push_int(&mut attrs, "targets", self.targets);
191        Ok(build_node::<Self>(attrs, inputs, outputs))
192    }
193
194    fn from_node(node: &Node) -> Result<Self, OpError> {
195        check_op::<Self>(node)?;
196        check_counts(Self::OP_TYPE, node, |n| n == 1, "1", |n| n == 1, "1")?;
197        let mut table = AttrTable::new(Self::OP_TYPE, node)?;
198        let coefficients = table.req_floats("coefficients")?;
199        let intercepts = table.opt_floats("intercepts")?;
200        let post_transform = table.opt_string("post_transform")?;
201        let targets = table.opt_int("targets")?;
202        table.finish()?;
203        Ok(LinearRegressor {
204            coefficients,
205            intercepts,
206            post_transform,
207            targets,
208        })
209    }
210}
211
212#[derive(Debug, Clone, PartialEq, Default)]
213pub struct SvmCommon {
214    pub coefficients: Option<Vec<f32>>,
215    pub kernel_params: Option<Vec<f32>>,
216    pub kernel_type: Option<String>,
217    pub post_transform: Option<String>,
218    pub prob_a: Option<Vec<f32>>,
219    pub prob_b: Option<Vec<f32>>,
220    pub rho: Option<Vec<f32>>,
221    pub support_vectors: Option<Vec<f32>>,
222}
223
224impl SvmCommon {
225    fn emit(&self, attrs: &mut Vec<crate::ir::Attribute>) {
226        push_floats(attrs, "coefficients", &self.coefficients);
227        push_floats(attrs, "kernel_params", &self.kernel_params);
228        push_string(attrs, "kernel_type", &self.kernel_type);
229        push_string(attrs, "post_transform", &self.post_transform);
230        push_floats(attrs, "prob_a", &self.prob_a);
231        push_floats(attrs, "prob_b", &self.prob_b);
232        push_floats(attrs, "rho", &self.rho);
233        push_floats(attrs, "support_vectors", &self.support_vectors);
234    }
235
236    fn parse(table: &mut AttrTable) -> Result<Self, OpError> {
237        Ok(SvmCommon {
238            coefficients: table.opt_floats("coefficients")?,
239            kernel_params: table.opt_floats("kernel_params")?,
240            kernel_type: table.opt_string("kernel_type")?,
241            post_transform: table.opt_string("post_transform")?,
242            prob_a: table.opt_floats("prob_a")?,
243            prob_b: table.opt_floats("prob_b")?,
244            rho: table.opt_floats("rho")?,
245            support_vectors: table.opt_floats("support_vectors")?,
246        })
247    }
248}
249
250#[derive(Debug, Clone, PartialEq)]
251pub struct SvmClassifier {
252    pub classlabels: ClassLabels,
253    pub common: SvmCommon,
254    pub vectors_per_class: Option<Vec<i64>>,
255}
256
257impl OnnxOp for SvmClassifier {
258    const OP_TYPE: &'static str = "SVMClassifier";
259    const DOMAIN: &'static str = ML_DOMAIN;
260
261    fn to_node(&self, inputs: Vec<String>, outputs: Vec<String>) -> Result<Emitted, OpError> {
262        self.classlabels.check(Self::OP_TYPE)?;
263        if inputs.len() != 1 {
264            return Err(OpError::WrongInputCount {
265                op: Self::OP_TYPE,
266                expected: "1".to_string(),
267                got: inputs.len(),
268            });
269        }
270        if outputs.len() != 2 {
271            return Err(OpError::WrongOutputCount {
272                op: Self::OP_TYPE,
273                expected: "2".to_string(),
274                got: outputs.len(),
275            });
276        }
277        let mut attrs = Vec::new();
278        self.classlabels.emit(&mut attrs);
279        self.common.emit(&mut attrs);
280        push_ints(&mut attrs, "vectors_per_class", &self.vectors_per_class);
281        Ok(build_node::<Self>(attrs, inputs, outputs))
282    }
283
284    fn from_node(node: &Node) -> Result<Self, OpError> {
285        check_op::<Self>(node)?;
286        check_counts(Self::OP_TYPE, node, |n| n == 1, "1", |n| n == 2, "2")?;
287        let mut table = AttrTable::new(Self::OP_TYPE, node)?;
288        let classlabels = ClassLabels::parse(Self::OP_TYPE, &mut table)?;
289        let common = SvmCommon::parse(&mut table)?;
290        let vectors_per_class = table.opt_ints("vectors_per_class")?;
291        table.finish()?;
292        Ok(SvmClassifier {
293            classlabels,
294            common,
295            vectors_per_class,
296        })
297    }
298}
299
300#[derive(Debug, Clone, PartialEq)]
301pub struct SvmRegressor {
302    pub common: SvmCommon,
303    pub n_supports: Option<i64>,
304    pub one_class: Option<i64>,
305}
306
307impl OnnxOp for SvmRegressor {
308    const OP_TYPE: &'static str = "SVMRegressor";
309    const DOMAIN: &'static str = ML_DOMAIN;
310
311    fn to_node(&self, inputs: Vec<String>, outputs: Vec<String>) -> Result<Emitted, OpError> {
312        if inputs.len() != 1 {
313            return Err(OpError::WrongInputCount {
314                op: Self::OP_TYPE,
315                expected: "1".to_string(),
316                got: inputs.len(),
317            });
318        }
319        if outputs.len() != 1 {
320            return Err(OpError::WrongOutputCount {
321                op: Self::OP_TYPE,
322                expected: "1".to_string(),
323                got: outputs.len(),
324            });
325        }
326        let mut attrs = Vec::new();
327        self.common.emit(&mut attrs);
328        push_int(&mut attrs, "n_supports", self.n_supports);
329        push_int(&mut attrs, "one_class", self.one_class);
330        Ok(build_node::<Self>(attrs, inputs, outputs))
331    }
332
333    fn from_node(node: &Node) -> Result<Self, OpError> {
334        check_op::<Self>(node)?;
335        check_counts(Self::OP_TYPE, node, |n| n == 1, "1", |n| n == 1, "1")?;
336        let mut table = AttrTable::new(Self::OP_TYPE, node)?;
337        let common = SvmCommon::parse(&mut table)?;
338        let n_supports = table.opt_int("n_supports")?;
339        let one_class = table.opt_int("one_class")?;
340        table.finish()?;
341        Ok(SvmRegressor {
342            common,
343            n_supports,
344            one_class,
345        })
346    }
347}