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}