1use std::cmp::Ordering;
4
5use nalgebra::{DMatrix, DVector};
6use runmat_builtins::{
7 BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinOutputMode,
8 BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor,
9 CellArray, CharArray, LogicalArray, ResolveContext, StringArray, StructValue, Tensor, Type,
10 Value,
11};
12use runmat_macros::runtime_builtin;
13
14use crate::builtins::common::tensor;
15use crate::{build_runtime_error, gather_if_needed_async, BuiltinResult, RuntimeError};
16
17const CLASSIFY_NAME: &str = "classify";
18const EPS: f64 = 1.0e-10;
19const REGULARIZATION: f64 = 1.0e-9;
20
21const OUTPUT_CLASS: BuiltinParamDescriptor = BuiltinParamDescriptor {
22 name: "class",
23 ty: BuiltinParamType::Any,
24 arity: BuiltinParamArity::Required,
25 default: None,
26 description: "Predicted class labels for sample observations.",
27};
28
29const OUTPUT_ERR: BuiltinParamDescriptor = BuiltinParamDescriptor {
30 name: "err",
31 ty: BuiltinParamType::NumericScalar,
32 arity: BuiltinParamArity::Optional,
33 default: None,
34 description: "Apparent training-set error rate.",
35};
36
37const OUTPUT_POSTERIOR: BuiltinParamDescriptor = BuiltinParamDescriptor {
38 name: "posterior",
39 ty: BuiltinParamType::NumericArray,
40 arity: BuiltinParamArity::Optional,
41 default: None,
42 description: "Posterior class probabilities for sample observations.",
43};
44
45const OUTPUT_LOGP: BuiltinParamDescriptor = BuiltinParamDescriptor {
46 name: "logp",
47 ty: BuiltinParamType::NumericArray,
48 arity: BuiltinParamArity::Optional,
49 default: None,
50 description: "Log unconditional density for sample observations.",
51};
52
53const OUTPUT_COEFF: BuiltinParamDescriptor = BuiltinParamDescriptor {
54 name: "coeff",
55 ty: BuiltinParamType::Any,
56 arity: BuiltinParamArity::Optional,
57 default: None,
58 description: "Pairwise discriminant boundary coefficients.",
59};
60
61const PARAM_SAMPLE: BuiltinParamDescriptor = BuiltinParamDescriptor {
62 name: "sample",
63 ty: BuiltinParamType::NumericArray,
64 arity: BuiltinParamArity::Required,
65 default: None,
66 description: "Sample observations in rows.",
67};
68
69const PARAM_TRAINING: BuiltinParamDescriptor = BuiltinParamDescriptor {
70 name: "training",
71 ty: BuiltinParamType::NumericArray,
72 arity: BuiltinParamArity::Required,
73 default: None,
74 description: "Training observations in rows.",
75};
76
77const PARAM_GROUP: BuiltinParamDescriptor = BuiltinParamDescriptor {
78 name: "group",
79 ty: BuiltinParamType::Any,
80 arity: BuiltinParamArity::Required,
81 default: None,
82 description: "Training group labels.",
83};
84
85const PARAM_TYPE: BuiltinParamDescriptor = BuiltinParamDescriptor {
86 name: "type",
87 ty: BuiltinParamType::StringScalar,
88 arity: BuiltinParamArity::Optional,
89 default: Some("linear"),
90 description: "Discriminant type: linear, quadratic, diagLinear, diagQuadratic, or mahalanobis.",
91};
92
93const PARAM_PRIOR: BuiltinParamDescriptor = BuiltinParamDescriptor {
94 name: "prior",
95 ty: BuiltinParamType::Any,
96 arity: BuiltinParamArity::Optional,
97 default: None,
98 description: "Prior probabilities as a numeric vector, 'empirical', or a struct with group and prob fields.",
99};
100
101const INPUTS_BASIC: [BuiltinParamDescriptor; 3] = [PARAM_SAMPLE, PARAM_TRAINING, PARAM_GROUP];
102const INPUTS_FULL: [BuiltinParamDescriptor; 5] = [
103 PARAM_SAMPLE,
104 PARAM_TRAINING,
105 PARAM_GROUP,
106 PARAM_TYPE,
107 PARAM_PRIOR,
108];
109const OUTPUT_CLASS_ONLY: [BuiltinParamDescriptor; 1] = [OUTPUT_CLASS];
110const OUTPUT_ALL: [BuiltinParamDescriptor; 5] = [
111 OUTPUT_CLASS,
112 OUTPUT_ERR,
113 OUTPUT_POSTERIOR,
114 OUTPUT_LOGP,
115 OUTPUT_COEFF,
116];
117
118const CLASSIFY_SIGNATURES: [BuiltinSignatureDescriptor; 3] = [
119 BuiltinSignatureDescriptor {
120 label: "class = classify(sample, training, group)",
121 inputs: &INPUTS_BASIC,
122 outputs: &OUTPUT_CLASS_ONLY,
123 },
124 BuiltinSignatureDescriptor {
125 label: "class = classify(sample, training, group, type, prior)",
126 inputs: &INPUTS_FULL,
127 outputs: &OUTPUT_CLASS_ONLY,
128 },
129 BuiltinSignatureDescriptor {
130 label: "[class,err,posterior,logp,coeff] = classify(___)",
131 inputs: &INPUTS_FULL,
132 outputs: &OUTPUT_ALL,
133 },
134];
135
136const ERROR_INVALID_ARGUMENT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
137 code: "RM.CLASSIFY.INVALID_ARGUMENT",
138 identifier: Some("RunMat:classify:InvalidArgument"),
139 when: "Inputs, labels, dimensions, discriminant type, or priors are malformed.",
140 message: "classify: invalid argument",
141};
142
143const ERROR_NUMERICAL: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
144 code: "RM.CLASSIFY.NUMERICAL",
145 identifier: Some("RunMat:classify:Numerical"),
146 when: "Covariance matrices cannot be solved numerically.",
147 message: "classify: numerical failure",
148};
149
150const ERROR_INTERNAL: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
151 code: "RM.CLASSIFY.INTERNAL",
152 identifier: Some("RunMat:classify:Internal"),
153 when: "RunMat cannot construct classification outputs.",
154 message: "classify: internal error",
155};
156
157const CLASSIFY_ERRORS: [BuiltinErrorDescriptor; 3] =
158 [ERROR_INVALID_ARGUMENT, ERROR_NUMERICAL, ERROR_INTERNAL];
159
160pub const CLASSIFY_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
161 signatures: &CLASSIFY_SIGNATURES,
162 output_mode: BuiltinOutputMode::ByRequestedOutputCount,
163 completion_policy: BuiltinCompletionPolicy::Public,
164 errors: &CLASSIFY_ERRORS,
165};
166
167fn classify_type(_args: &[Type], _ctx: &ResolveContext) -> Type {
168 Type::Unknown
169}
170
171fn classify_error(
172 message: impl Into<String>,
173 descriptor: &'static BuiltinErrorDescriptor,
174) -> RuntimeError {
175 let mut builder = build_runtime_error(message).with_builtin(CLASSIFY_NAME);
176 if let Some(identifier) = descriptor.identifier {
177 builder = builder.with_identifier(identifier);
178 }
179 builder.build()
180}
181
182fn invalid_argument(message: impl Into<String>) -> RuntimeError {
183 classify_error(message, &ERROR_INVALID_ARGUMENT)
184}
185
186fn numerical_error(message: impl Into<String>) -> RuntimeError {
187 classify_error(message, &ERROR_NUMERICAL)
188}
189
190fn internal_error(message: impl Into<String>) -> RuntimeError {
191 classify_error(message, &ERROR_INTERNAL)
192}
193
194#[derive(Clone, Copy, Debug, Eq, PartialEq)]
195enum DiscriminantType {
196 Linear,
197 Quadratic,
198 DiagLinear,
199 DiagQuadratic,
200 Mahalanobis,
201}
202
203impl DiscriminantType {
204 fn property_name(self) -> &'static str {
205 match self {
206 Self::Linear => "linear",
207 Self::Quadratic => "quadratic",
208 Self::DiagLinear => "diagLinear",
209 Self::DiagQuadratic => "diagQuadratic",
210 Self::Mahalanobis => "mahalanobis",
211 }
212 }
213
214 fn uses_shared_covariance(self) -> bool {
215 matches!(self, Self::Linear | Self::DiagLinear)
216 }
217
218 fn uses_diagonal_covariance(self) -> bool {
219 matches!(self, Self::DiagLinear | Self::DiagQuadratic)
220 }
221}
222
223#[derive(Clone, Copy, Debug, Eq, PartialEq)]
224enum LabelKind {
225 Numeric,
226 String,
227 Char,
228 Cell,
229 Logical,
230}
231
232#[derive(Clone, Debug, PartialEq)]
233enum ClassLabel {
234 Numeric(f64),
235 Text(String),
236 Logical(bool),
237}
238
239#[derive(Clone, Debug)]
240struct LabelVector {
241 labels: Vec<ClassLabel>,
242 kind: LabelKind,
243}
244
245#[derive(Clone, Debug)]
246struct ClassStats {
247 label: ClassLabel,
248 rows: Vec<usize>,
249 mean: Vec<f64>,
250 covariance: DMatrix<f64>,
251 inverse: DMatrix<f64>,
252 log_det: f64,
253}
254
255#[derive(Clone, Debug)]
256struct PreparedModel {
257 kind: LabelKind,
258 discriminant: DiscriminantType,
259 classes: Vec<ClassStats>,
260 priors: Vec<f64>,
261 shared_inverse: Option<DMatrix<f64>>,
262 shared_log_det: Option<f64>,
263 cols: usize,
264}
265
266#[derive(Clone, Debug)]
267struct PredictionResult {
268 labels: Vec<usize>,
269 posterior: Tensor,
270 logp: Tensor,
271}
272
273#[runtime_builtin(
274 name = "classify",
275 category = "stats/ml",
276 summary = "Classify observations using discriminant analysis.",
277 keywords = "classify,discriminant analysis,lda,qda,classification,statistics,machine learning",
278 type_resolver(classify_type),
279 descriptor(crate::builtins::stats::ml::classification::CLASSIFY_DESCRIPTOR),
280 builtin_path = "crate::builtins::stats::ml::classification"
281)]
282async fn classify_builtin(
283 sample: Value,
284 training: Value,
285 group: Value,
286 rest: Vec<Value>,
287) -> BuiltinResult<Value> {
288 let requested_outputs = crate::output_count::current_output_count();
289 if let Some(0) = requested_outputs {
290 return Ok(Value::OutputList(Vec::new()));
291 }
292 if matches!(requested_outputs, Some(count) if count > 5) {
293 return Err(invalid_argument("classify: too many output arguments"));
294 }
295 let sample = gathered(sample).await?;
296 let training = gathered(training).await?;
297 let group = gathered(group).await?;
298 let rest = gather_values(rest).await?;
299 let output_limit = requested_outputs.unwrap_or(1).max(1);
300 let result = classify_compute(sample, training, group, rest, output_limit)?;
301 match requested_outputs {
302 Some(1) => Ok(Value::OutputList(vec![result[0].clone()])),
303 Some(out_count) => Ok(crate::output_count::output_list_with_padding(
304 out_count, result,
305 )),
306 None => Ok(result[0].clone()),
307 }
308}
309
310async fn gathered(value: Value) -> BuiltinResult<Value> {
311 gather_if_needed_async(&value)
312 .await
313 .map_err(|err| invalid_argument(format!("classify: {err}")))
314}
315
316async fn gather_values(values: Vec<Value>) -> BuiltinResult<Vec<Value>> {
317 let mut out = Vec::with_capacity(values.len());
318 for value in values {
319 out.push(gathered(value).await?);
320 }
321 Ok(out)
322}
323
324fn classify_compute(
325 sample: Value,
326 training: Value,
327 group: Value,
328 rest: Vec<Value>,
329 output_limit: usize,
330) -> BuiltinResult<Vec<Value>> {
331 if rest.len() > 2 {
332 return Err(invalid_argument(
333 "classify: accepts at most type and prior arguments",
334 ));
335 }
336 let sample = numeric_matrix(sample, "sample")?;
337 let training = numeric_matrix(training, "training")?;
338 if sample.cols != training.cols {
339 return Err(invalid_argument(
340 "classify: sample and training must have the same number of columns",
341 ));
342 }
343 if sample.cols == 0 || training.rows == 0 {
344 return Err(invalid_argument(
345 "classify: sample and training must be nonempty numeric matrices",
346 ));
347 }
348 if sample.data.iter().any(|value| !value.is_finite())
349 || training.data.iter().any(|value| !value.is_finite())
350 {
351 return Err(invalid_argument(
352 "classify: sample and training must contain finite values",
353 ));
354 }
355 let labels = labels_from_value(group, "group")?;
356 if labels.labels.len() != training.rows {
357 return Err(invalid_argument(
358 "classify: group length must match the number of rows in training",
359 ));
360 }
361 let discriminant = if let Some(value) = rest.first() {
362 parse_type(value)?
363 } else {
364 DiscriminantType::Linear
365 };
366 let prepared = prepare_model(&training, labels, discriminant, rest.get(1))?;
367 let predictions = predict_rows(&prepared, &sample)?;
368 let class = labels_value(&prepared.classes, prepared.kind, &predictions.labels)?;
369 if output_limit == 1 {
370 return Ok(vec![class]);
371 }
372 let training_predictions = predict_rows(&prepared, &training)?;
373 let err = apparent_error(&prepared, &training_predictions.labels);
374 let mut outputs = vec![class, Value::Num(err)];
375 if output_limit >= 3 {
376 outputs.push(Value::Tensor(predictions.posterior));
377 }
378 if output_limit >= 4 {
379 outputs.push(Value::Tensor(predictions.logp));
380 }
381 if output_limit >= 5 {
382 outputs.push(coeff_value(&prepared)?);
383 }
384 Ok(outputs)
385}
386
387fn numeric_matrix(value: Value, name: &str) -> BuiltinResult<Tensor> {
388 let tensor = tensor::value_into_tensor_for(CLASSIFY_NAME, value)
389 .map_err(|err| invalid_argument(format!("classify: {name}: {err}")))?;
390 if tensor.shape.len() > 2 {
391 return Err(invalid_argument(format!(
392 "classify: {name} must be a 2-D numeric matrix"
393 )));
394 }
395 Ok(tensor)
396}
397
398fn parse_type(value: &Value) -> BuiltinResult<DiscriminantType> {
399 let text = scalar_text(value, "type")?;
400 match canonical(&text).as_str() {
401 "linear" => Ok(DiscriminantType::Linear),
402 "quadratic" => Ok(DiscriminantType::Quadratic),
403 "diaglinear" => Ok(DiscriminantType::DiagLinear),
404 "diagquadratic" => Ok(DiscriminantType::DiagQuadratic),
405 "mahalanobis" => Ok(DiscriminantType::Mahalanobis),
406 other => Err(invalid_argument(format!(
407 "classify: unsupported discriminant type '{other}'"
408 ))),
409 }
410}
411
412fn prepare_model(
413 training: &Tensor,
414 labels: LabelVector,
415 discriminant: DiscriminantType,
416 prior_value: Option<&Value>,
417) -> BuiltinResult<PreparedModel> {
418 let unique = unique_labels(&labels.labels, labels.kind);
419 if unique.len() < 2 {
420 return Err(invalid_argument(
421 "classify: group must contain at least two nonmissing classes",
422 ));
423 }
424 let mut rows_by_class = vec![Vec::<usize>::new(); unique.len()];
425 for (row, label) in labels.labels.iter().enumerate() {
426 if is_missing_label(label) {
427 continue;
428 }
429 let Some(class_index) = unique.iter().position(|class| same_label(class, label)) else {
430 continue;
431 };
432 rows_by_class[class_index].push(row);
433 }
434 if rows_by_class.iter().any(Vec::is_empty) {
435 return Err(invalid_argument(
436 "classify: every class must contain at least one training row",
437 ));
438 }
439 let priors = parse_prior(prior_value, &unique, labels.kind, &rows_by_class)?;
440 let cols = training.cols;
441 let mut classes = Vec::with_capacity(unique.len());
442 for (class_index, rows) in rows_by_class.iter().enumerate() {
443 let mean = class_mean(training, rows);
444 let mut covariance = class_covariance(training, rows, &mean)?;
445 if discriminant.uses_diagonal_covariance() {
446 covariance = diagonalized(&covariance);
447 }
448 let (inverse, log_det) = invert_covariance(&covariance, "class covariance")?;
449 classes.push(ClassStats {
450 label: unique[class_index].clone(),
451 rows: rows.clone(),
452 mean,
453 covariance,
454 inverse,
455 log_det,
456 });
457 }
458 let (shared_inverse, shared_log_det) = if discriminant.uses_shared_covariance() {
459 let mut pooled = pooled_covariance(&classes, cols)?;
460 if discriminant.uses_diagonal_covariance() {
461 pooled = diagonalized(&pooled);
462 }
463 let (inverse, log_det) = invert_covariance(&pooled, "pooled covariance")?;
464 (Some(inverse), Some(log_det))
465 } else {
466 (None, None)
467 };
468 Ok(PreparedModel {
469 kind: labels.kind,
470 discriminant,
471 classes,
472 priors,
473 shared_inverse,
474 shared_log_det,
475 cols,
476 })
477}
478
479fn class_mean(training: &Tensor, rows: &[usize]) -> Vec<f64> {
480 let mut mean = vec![0.0; training.cols];
481 for row in rows {
482 for col in 0..training.cols {
483 mean[col] += training.data[row + col * training.rows];
484 }
485 }
486 for value in &mut mean {
487 *value /= rows.len() as f64;
488 }
489 mean
490}
491
492fn class_covariance(
493 training: &Tensor,
494 rows: &[usize],
495 mean: &[f64],
496) -> BuiltinResult<DMatrix<f64>> {
497 let cols = training.cols;
498 let mut cov = DMatrix::<f64>::zeros(cols, cols);
499 if rows.len() < 2 {
500 for idx in 0..cols {
501 cov[(idx, idx)] = REGULARIZATION;
502 }
503 return Ok(cov);
504 }
505 for row in rows {
506 for a in 0..cols {
507 let da = training.data[row + a * training.rows] - mean[a];
508 for b in 0..cols {
509 let db = training.data[row + b * training.rows] - mean[b];
510 cov[(a, b)] += da * db;
511 }
512 }
513 }
514 let denom = (rows.len() - 1) as f64;
515 for value in cov.iter_mut() {
516 *value /= denom;
517 }
518 Ok(cov)
519}
520
521fn pooled_covariance(classes: &[ClassStats], cols: usize) -> BuiltinResult<DMatrix<f64>> {
522 let total_rows = classes.iter().map(|class| class.rows.len()).sum::<usize>();
523 if total_rows <= classes.len() {
524 let mut identity = DMatrix::<f64>::zeros(cols, cols);
525 for idx in 0..cols {
526 identity[(idx, idx)] = REGULARIZATION;
527 }
528 return Ok(identity);
529 }
530 let mut pooled = DMatrix::<f64>::zeros(cols, cols);
531 for class in classes {
532 let weight = class.rows.len().saturating_sub(1) as f64;
533 pooled += class.covariance.clone() * weight;
534 }
535 pooled /= (total_rows - classes.len()) as f64;
536 Ok(pooled)
537}
538
539fn diagonalized(matrix: &DMatrix<f64>) -> DMatrix<f64> {
540 let mut out = DMatrix::<f64>::zeros(matrix.nrows(), matrix.ncols());
541 for idx in 0..matrix.nrows().min(matrix.ncols()) {
542 out[(idx, idx)] = matrix[(idx, idx)];
543 }
544 out
545}
546
547fn invert_covariance(
548 matrix: &DMatrix<f64>,
549 label: &'static str,
550) -> BuiltinResult<(DMatrix<f64>, f64)> {
551 let mut adjusted = matrix.clone();
552 for idx in 0..adjusted.nrows().min(adjusted.ncols()) {
553 adjusted[(idx, idx)] += REGULARIZATION;
554 }
555 let lu = adjusted.lu();
556 let det = lu.determinant();
557 if !det.is_finite() || det.abs() <= EPS {
558 return Err(numerical_error(format!("classify: singular {label}")));
559 }
560 let Some(inverse) = lu.try_inverse() else {
561 return Err(numerical_error(format!("classify: singular {label}")));
562 };
563 Ok((inverse, det.abs().ln()))
564}
565
566fn predict_rows(model: &PreparedModel, data: &Tensor) -> BuiltinResult<PredictionResult> {
567 let rows = data.rows;
568 let class_count = model.classes.len();
569 let mut label_indices = Vec::with_capacity(rows);
570 let mut posterior_data = vec![0.0; rows * class_count];
571 let mut logp_data = vec![0.0; rows];
572 for row in 0..rows {
573 let x = row_vector(data, row);
574 let mut scores = Vec::with_capacity(class_count);
575 for class_index in 0..class_count {
576 scores.push(discriminant_score(model, class_index, &x)?);
577 }
578 let best = scores
579 .iter()
580 .enumerate()
581 .max_by(|(_, left), (_, right)| left.partial_cmp(right).unwrap_or(Ordering::Equal))
582 .map(|(idx, _)| idx)
583 .ok_or_else(|| internal_error("classify: no classes"))?;
584 label_indices.push(best);
585 if model.discriminant == DiscriminantType::Mahalanobis {
586 for class_index in 0..class_count {
587 posterior_data[row + class_index * rows] = f64::NAN;
588 }
589 logp_data[row] = f64::NAN;
590 } else {
591 let logp = log_sum_exp(&scores);
592 logp_data[row] = logp;
593 for class_index in 0..class_count {
594 posterior_data[row + class_index * rows] = (scores[class_index] - logp).exp();
595 }
596 }
597 }
598 Ok(PredictionResult {
599 labels: label_indices,
600 posterior: Tensor::new(posterior_data, vec![rows, class_count])
601 .map_err(|err| internal_error(format!("classify: {err}")))?,
602 logp: Tensor::new(logp_data, vec![rows, 1])
603 .map_err(|err| internal_error(format!("classify: {err}")))?,
604 })
605}
606
607fn row_vector(data: &Tensor, row: usize) -> DVector<f64> {
608 DVector::from_iterator(
609 data.cols,
610 (0..data.cols).map(|col| data.data[row + col * data.rows]),
611 )
612}
613
614fn mean_vector(mean: &[f64]) -> DVector<f64> {
615 DVector::from_iterator(mean.len(), mean.iter().copied())
616}
617
618fn discriminant_score(
619 model: &PreparedModel,
620 class_index: usize,
621 x: &DVector<f64>,
622) -> BuiltinResult<f64> {
623 let class = &model.classes[class_index];
624 let mean = mean_vector(&class.mean);
625 let diff = x - &mean;
626 let (inverse, log_det) = if let Some(inverse) = &model.shared_inverse {
627 (inverse, model.shared_log_det.unwrap_or(0.0))
628 } else {
629 (&class.inverse, class.log_det)
630 };
631 let distance = (diff.transpose() * inverse * diff)[(0, 0)];
632 if model.discriminant == DiscriminantType::Mahalanobis {
633 Ok(-distance)
634 } else {
635 let gaussian_norm = -0.5 * (model.cols as f64) * (2.0 * std::f64::consts::PI).ln();
636 Ok(model.priors[class_index].ln() + gaussian_norm - 0.5 * log_det - 0.5 * distance)
637 }
638}
639
640fn log_sum_exp(values: &[f64]) -> f64 {
641 let max = values
642 .iter()
643 .copied()
644 .fold(f64::NEG_INFINITY, |acc, value| acc.max(value));
645 max + values
646 .iter()
647 .map(|value| (value - max).exp())
648 .sum::<f64>()
649 .ln()
650}
651
652fn apparent_error(model: &PreparedModel, predicted: &[usize]) -> f64 {
653 let mut total = 0.0;
654 for (class_index, class) in model.classes.iter().enumerate() {
655 let mut wrong = 0usize;
656 for row in &class.rows {
657 if predicted.get(*row).copied() != Some(class_index) {
658 wrong += 1;
659 }
660 }
661 let class_error = if class.rows.is_empty() {
662 0.0
663 } else {
664 wrong as f64 / class.rows.len() as f64
665 };
666 total += model.priors[class_index] * class_error;
667 }
668 total
669}
670
671fn parse_prior(
672 prior: Option<&Value>,
673 labels: &[ClassLabel],
674 kind: LabelKind,
675 rows_by_class: &[Vec<usize>],
676) -> BuiltinResult<Vec<f64>> {
677 let raw = match prior {
678 None => vec![1.0; labels.len()],
679 Some(value) if string_matches(value, "empirical") => rows_by_class
680 .iter()
681 .map(|rows| rows.len() as f64)
682 .collect::<Vec<_>>(),
683 Some(Value::Struct(st)) => prior_from_struct(st, labels, kind)?,
684 Some(value) => numeric_vector(value, "prior")?,
685 };
686 if raw.len() != labels.len() {
687 return Err(invalid_argument(
688 "classify: prior must contain one probability per class",
689 ));
690 }
691 if raw.iter().any(|value| !value.is_finite() || *value < 0.0) {
692 return Err(invalid_argument(
693 "classify: prior probabilities must be finite and nonnegative",
694 ));
695 }
696 let total = raw.iter().sum::<f64>();
697 if total <= 0.0 {
698 return Err(invalid_argument(
699 "classify: prior probabilities must have positive total",
700 ));
701 }
702 Ok(raw.into_iter().map(|value| value / total).collect())
703}
704
705fn prior_from_struct(
706 st: &StructValue,
707 labels: &[ClassLabel],
708 kind: LabelKind,
709) -> BuiltinResult<Vec<f64>> {
710 let group_value = st
711 .fields
712 .get("group")
713 .or_else(|| st.fields.get("Group"))
714 .ok_or_else(|| invalid_argument("classify: prior struct must contain group field"))?
715 .clone();
716 let prob_value = st
717 .fields
718 .get("prob")
719 .or_else(|| st.fields.get("Prob"))
720 .ok_or_else(|| invalid_argument("classify: prior struct must contain prob field"))?;
721 let groups = labels_from_value(group_value, "prior.group")?;
722 if !label_kinds_compatible(groups.kind, kind) {
723 return Err(invalid_argument(
724 "classify: prior group labels must match group label type",
725 ));
726 }
727 let probs = numeric_vector(prob_value, "prior.prob")?;
728 if groups.labels.len() != probs.len() {
729 return Err(invalid_argument(
730 "classify: prior group and prob lengths must match",
731 ));
732 }
733 let mut out = vec![0.0; labels.len()];
734 for (group, prob) in groups.labels.iter().zip(probs.iter()) {
735 let Some(index) = labels.iter().position(|label| same_label(label, group)) else {
736 return Err(invalid_argument(
737 "classify: prior struct contains unknown group label",
738 ));
739 };
740 out[index] = *prob;
741 }
742 Ok(out)
743}
744
745fn coeff_value(model: &PreparedModel) -> BuiltinResult<Value> {
746 let k = model.classes.len();
747 let mut entries = Vec::with_capacity(k * k);
748 for row in 0..k {
749 for col in 0..k {
750 let coeff = pairwise_coeff(model, row, col)?;
751 let mut st = StructValue::new();
752 st.insert(
753 "type",
754 Value::String(model.discriminant.property_name().to_string()),
755 );
756 st.insert("name1", labels_value(&model.classes, model.kind, &[row])?);
757 st.insert("name2", labels_value(&model.classes, model.kind, &[col])?);
758 st.insert("const", Value::Num(coeff.0));
759 st.insert(
760 "linear",
761 Value::Tensor(
762 Tensor::new(coeff.1, vec![model.cols, 1])
763 .map_err(|err| internal_error(format!("classify: {err}")))?,
764 ),
765 );
766 if !matches!(
767 model.discriminant,
768 DiscriminantType::Linear | DiscriminantType::DiagLinear
769 ) {
770 st.insert(
771 "quadratic",
772 Value::Tensor(
773 Tensor::new(coeff.2, vec![model.cols, model.cols])
774 .map_err(|err| internal_error(format!("classify: {err}")))?,
775 ),
776 );
777 }
778 entries.push(Value::Struct(st));
779 }
780 }
781 Ok(Value::Cell(CellArray::new(entries, k, k).map_err(
782 |err| internal_error(format!("classify: {err}")),
783 )?))
784}
785
786fn pairwise_coeff(
787 model: &PreparedModel,
788 i: usize,
789 j: usize,
790) -> BuiltinResult<(f64, Vec<f64>, Vec<f64>)> {
791 if i == j {
792 return Ok((
793 0.0,
794 vec![0.0; model.cols],
795 vec![0.0; model.cols * model.cols],
796 ));
797 }
798 let class_i = &model.classes[i];
799 let class_j = &model.classes[j];
800 let mean_i = mean_vector(&class_i.mean);
801 let mean_j = mean_vector(&class_j.mean);
802 let (inv_i, logdet_i) = if let Some(inv) = &model.shared_inverse {
803 (inv, model.shared_log_det.unwrap_or(0.0))
804 } else {
805 (&class_i.inverse, class_i.log_det)
806 };
807 let (inv_j, logdet_j) = if let Some(inv) = &model.shared_inverse {
808 (inv, model.shared_log_det.unwrap_or(0.0))
809 } else {
810 (&class_j.inverse, class_j.log_det)
811 };
812 let linear_vec = mean_i.transpose() * inv_i - mean_j.transpose() * inv_j;
813 let const_term = if model.discriminant == DiscriminantType::Mahalanobis {
814 -0.5 * ((mean_i.transpose() * inv_i * &mean_i)[(0, 0)]
815 - (mean_j.transpose() * inv_j * &mean_j)[(0, 0)])
816 } else {
817 log_prior_ratio(model.priors[i], model.priors[j])
818 - 0.5 * (logdet_i - logdet_j)
819 - 0.5
820 * ((mean_i.transpose() * inv_i * &mean_i)[(0, 0)]
821 - (mean_j.transpose() * inv_j * &mean_j)[(0, 0)])
822 };
823 let quadratic = if model.discriminant.uses_shared_covariance() {
824 DMatrix::<f64>::zeros(model.cols, model.cols)
825 } else {
826 (inv_j - inv_i) * 0.5
827 };
828 Ok((
829 const_term,
830 linear_vec.iter().copied().collect(),
831 quadratic.iter().copied().collect(),
832 ))
833}
834
835fn log_prior_ratio(left: f64, right: f64) -> f64 {
836 match (left > 0.0, right > 0.0) {
837 (true, true) => (left / right).ln(),
838 (true, false) => f64::INFINITY,
839 (false, true) => f64::NEG_INFINITY,
840 (false, false) => 0.0,
841 }
842}
843
844fn labels_from_value(value: Value, name: &str) -> BuiltinResult<LabelVector> {
845 match value {
846 Value::Tensor(tensor) => Ok(LabelVector {
847 labels: vector_values(&tensor, name)?
848 .into_iter()
849 .map(ClassLabel::Numeric)
850 .collect(),
851 kind: LabelKind::Numeric,
852 }),
853 Value::Num(value) => Ok(LabelVector {
854 labels: vec![ClassLabel::Numeric(value)],
855 kind: LabelKind::Numeric,
856 }),
857 Value::Int(value) => Ok(LabelVector {
858 labels: vec![ClassLabel::Numeric(value.to_f64())],
859 kind: LabelKind::Numeric,
860 }),
861 Value::String(text) => Ok(LabelVector {
862 labels: vec![ClassLabel::Text(text)],
863 kind: LabelKind::String,
864 }),
865 Value::CharArray(array) => Ok(LabelVector {
866 labels: char_rows(&array).into_iter().map(ClassLabel::Text).collect(),
867 kind: LabelKind::Char,
868 }),
869 Value::StringArray(array) => Ok(LabelVector {
870 labels: array.data.into_iter().map(ClassLabel::Text).collect(),
871 kind: LabelKind::String,
872 }),
873 Value::Cell(cell) => {
874 let mut labels = Vec::with_capacity(cell.data.len());
875 for item in cell.data {
876 labels.push(ClassLabel::Text(scalar_text(&item, name)?));
877 }
878 Ok(LabelVector {
879 labels,
880 kind: LabelKind::Cell,
881 })
882 }
883 Value::Bool(value) => Ok(LabelVector {
884 labels: vec![ClassLabel::Logical(value)],
885 kind: LabelKind::Logical,
886 }),
887 Value::LogicalArray(array) => Ok(LabelVector {
888 labels: array
889 .data
890 .into_iter()
891 .map(|value| ClassLabel::Logical(value != 0))
892 .collect(),
893 kind: LabelKind::Logical,
894 }),
895 other => Err(invalid_argument(format!(
896 "classify: {name} must be numeric, logical, string, character, or cellstr labels; got {other:?}"
897 ))),
898 }
899}
900
901fn vector_values(tensor: &Tensor, name: &str) -> BuiltinResult<Vec<f64>> {
902 if tensor.shape.iter().filter(|dim| **dim > 1).count() > 1 {
903 return Err(invalid_argument(format!(
904 "classify: {name} must be a vector"
905 )));
906 }
907 Ok(tensor.data.clone())
908}
909
910fn numeric_vector(value: &Value, name: &str) -> BuiltinResult<Vec<f64>> {
911 match value {
912 Value::Tensor(tensor) => vector_values(tensor, name),
913 Value::Num(value) => Ok(vec![*value]),
914 Value::Int(value) => Ok(vec![value.to_f64()]),
915 other => Err(invalid_argument(format!(
916 "classify: {name} must be a numeric vector, got {other:?}"
917 ))),
918 }
919}
920
921fn unique_labels(labels: &[ClassLabel], kind: LabelKind) -> Vec<ClassLabel> {
922 let mut out = Vec::new();
923 for label in labels {
924 if is_missing_label(label) {
925 continue;
926 }
927 if !out.iter().any(|existing| same_label(existing, label)) {
928 out.push(label.clone());
929 }
930 }
931 match kind {
932 LabelKind::Numeric => out.sort_by(|left, right| match (left, right) {
933 (ClassLabel::Numeric(a), ClassLabel::Numeric(b)) => {
934 a.partial_cmp(b).unwrap_or(Ordering::Equal)
935 }
936 _ => Ordering::Equal,
937 }),
938 LabelKind::String | LabelKind::Char | LabelKind::Cell => {
939 out.sort_by(|left, right| match (left, right) {
940 (ClassLabel::Text(a), ClassLabel::Text(b)) => a.cmp(b),
941 _ => Ordering::Equal,
942 })
943 }
944 LabelKind::Logical => out.sort_by_key(|label| match label {
945 ClassLabel::Logical(value) => *value,
946 _ => false,
947 }),
948 }
949 out
950}
951
952fn same_label(left: &ClassLabel, right: &ClassLabel) -> bool {
953 match (left, right) {
954 (ClassLabel::Numeric(a), ClassLabel::Numeric(b)) => (a == b) || (a.is_nan() && b.is_nan()),
955 (ClassLabel::Text(a), ClassLabel::Text(b)) => a == b,
956 (ClassLabel::Logical(a), ClassLabel::Logical(b)) => a == b,
957 _ => false,
958 }
959}
960
961fn is_missing_label(label: &ClassLabel) -> bool {
962 match label {
963 ClassLabel::Numeric(value) => value.is_nan(),
964 ClassLabel::Text(value) => value.is_empty() || value == "<missing>",
965 ClassLabel::Logical(_) => false,
966 }
967}
968
969fn label_kinds_compatible(left: LabelKind, right: LabelKind) -> bool {
970 left == right || (is_text_kind(left) && is_text_kind(right))
971}
972
973fn is_text_kind(kind: LabelKind) -> bool {
974 matches!(kind, LabelKind::String | LabelKind::Char | LabelKind::Cell)
975}
976
977fn labels_value(
978 classes: &[ClassStats],
979 kind: LabelKind,
980 indices: &[usize],
981) -> BuiltinResult<Value> {
982 let labels = classes
983 .iter()
984 .map(|class| class.label.clone())
985 .collect::<Vec<_>>();
986 labels_from_indices(&labels, kind, indices)
987}
988
989fn labels_from_indices(
990 labels: &[ClassLabel],
991 kind: LabelKind,
992 indices: &[usize],
993) -> BuiltinResult<Value> {
994 match kind {
995 LabelKind::Numeric => {
996 let data = indices
997 .iter()
998 .map(|idx| match labels.get(*idx) {
999 Some(ClassLabel::Numeric(value)) => *value,
1000 _ => f64::NAN,
1001 })
1002 .collect::<Vec<_>>();
1003 Ok(Value::Tensor(
1004 Tensor::new(data, vec![indices.len(), 1])
1005 .map_err(|err| internal_error(format!("classify: {err}")))?,
1006 ))
1007 }
1008 LabelKind::String => {
1009 let data = indices
1010 .iter()
1011 .map(|idx| match labels.get(*idx) {
1012 Some(ClassLabel::Text(value)) => value.clone(),
1013 _ => String::new(),
1014 })
1015 .collect::<Vec<_>>();
1016 Ok(Value::StringArray(
1017 StringArray::new(data, vec![indices.len(), 1])
1018 .map_err(|err| internal_error(format!("classify: {err}")))?,
1019 ))
1020 }
1021 LabelKind::Char => {
1022 let rows = indices
1023 .iter()
1024 .map(|idx| match labels.get(*idx) {
1025 Some(ClassLabel::Text(value)) => value.clone(),
1026 _ => String::new(),
1027 })
1028 .collect::<Vec<_>>();
1029 Ok(Value::CharArray(char_array_from_rows(&rows)?))
1030 }
1031 LabelKind::Cell => {
1032 let data = indices
1033 .iter()
1034 .map(|idx| match labels.get(*idx) {
1035 Some(ClassLabel::Text(value)) => Value::CharArray(CharArray::new_row(value)),
1036 _ => Value::CharArray(CharArray::new_row("")),
1037 })
1038 .collect::<Vec<_>>();
1039 Ok(Value::Cell(
1040 CellArray::new(data, indices.len(), 1)
1041 .map_err(|err| internal_error(format!("classify: {err}")))?,
1042 ))
1043 }
1044 LabelKind::Logical => {
1045 let data = indices
1046 .iter()
1047 .map(|idx| match labels.get(*idx) {
1048 Some(ClassLabel::Logical(value)) if *value => 1,
1049 _ => 0,
1050 })
1051 .collect::<Vec<_>>();
1052 Ok(Value::LogicalArray(
1053 LogicalArray::new(data, vec![indices.len(), 1])
1054 .map_err(|err| internal_error(format!("classify: {err}")))?,
1055 ))
1056 }
1057 }
1058}
1059
1060fn char_rows(array: &CharArray) -> Vec<String> {
1061 let mut rows = Vec::with_capacity(array.rows);
1062 for row in 0..array.rows {
1063 let start = row * array.cols;
1064 rows.push(array.data[start..start + array.cols].iter().collect());
1065 }
1066 rows
1067}
1068
1069fn char_array_from_rows(rows: &[String]) -> BuiltinResult<CharArray> {
1070 let width = rows
1071 .iter()
1072 .map(|row| row.chars().count())
1073 .max()
1074 .unwrap_or(0);
1075 let mut data = Vec::with_capacity(rows.len() * width);
1076 for row in rows {
1077 let mut chars = row.chars().collect::<Vec<_>>();
1078 chars.resize(width, ' ');
1079 data.extend(chars);
1080 }
1081 CharArray::new(data, rows.len(), width)
1082 .map_err(|err| internal_error(format!("classify: {err}")))
1083}
1084
1085fn scalar_text(value: &Value, name: &str) -> BuiltinResult<String> {
1086 match value {
1087 Value::String(text) => Ok(text.clone()),
1088 Value::CharArray(chars) => Ok(chars.data.iter().collect()),
1089 Value::StringArray(array) if array.data.len() == 1 => Ok(array.data[0].clone()),
1090 other => Err(invalid_argument(format!(
1091 "classify: {name} must be a string scalar, got {other:?}"
1092 ))),
1093 }
1094}
1095
1096fn string_matches(value: &Value, expected: &str) -> bool {
1097 scalar_text(value, "text")
1098 .map(|text| text.eq_ignore_ascii_case(expected))
1099 .unwrap_or(false)
1100}
1101
1102fn canonical(text: &str) -> String {
1103 text.chars()
1104 .filter(|ch| *ch != '_' && *ch != '-' && !ch.is_whitespace())
1105 .collect::<String>()
1106 .to_ascii_lowercase()
1107}
1108
1109#[cfg(test)]
1110mod tests {
1111 use super::*;
1112 use futures::executor::block_on;
1113
1114 fn tensor(data: Vec<f64>, shape: Vec<usize>) -> Value {
1115 Value::Tensor(Tensor::new(data, shape).unwrap())
1116 }
1117
1118 fn classify(sample: Value, training: Value, group: Value, rest: Vec<Value>) -> Vec<Value> {
1119 let Value::OutputList(values) =
1120 block_on(classify_builtin(sample, training, group, rest)).expect("classify")
1121 else {
1122 panic!("expected output list");
1123 };
1124 values
1125 }
1126
1127 fn with_outputs<T>(count: usize, f: impl FnOnce() -> T) -> T {
1128 let _guard = crate::output_count::push_output_count(Some(count));
1129 f()
1130 }
1131
1132 #[test]
1133 fn classify_linear_predicts_numeric_groups() {
1134 with_outputs(5, || {
1135 let values = classify(
1136 tensor(vec![0.2, 2.2], vec![2, 1]),
1137 tensor(vec![0.0, 0.5, 2.0, 2.5], vec![4, 1]),
1138 tensor(vec![1.0, 1.0, 2.0, 2.0], vec![4, 1]),
1139 Vec::new(),
1140 );
1141 let Value::Tensor(labels) = &values[0] else {
1142 panic!("labels");
1143 };
1144 assert_eq!(labels.data, vec![1.0, 2.0]);
1145 assert!(matches!(values[1], Value::Num(err) if err <= EPS));
1146 let Value::Tensor(posterior) = &values[2] else {
1147 panic!("posterior");
1148 };
1149 assert_eq!(posterior.shape, vec![2, 2]);
1150 assert!(posterior.data[0] > posterior.data[2]);
1151 assert!(posterior.data[3] > posterior.data[1]);
1152 let Value::Cell(coeff) = &values[4] else {
1153 panic!("coeff");
1154 };
1155 assert_eq!(coeff.rows, 2);
1156 assert_eq!(coeff.cols, 2);
1157 assert!(
1158 matches!(&coeff.data[1], Value::Struct(st) if st.fields.contains_key("linear"))
1159 );
1160 });
1161 }
1162
1163 #[test]
1164 fn classify_quadratic_preserves_text_labels() {
1165 with_outputs(5, || {
1166 let groups = Value::StringArray(
1167 StringArray::new(
1168 vec!["low".into(), "low".into(), "high".into(), "high".into()],
1169 vec![4, 1],
1170 )
1171 .unwrap(),
1172 );
1173 let values = classify(
1174 tensor(vec![0.1, 3.2], vec![2, 1]),
1175 tensor(vec![0.0, 0.4, 3.0, 3.4], vec![4, 1]),
1176 groups,
1177 vec![Value::String("quadratic".into())],
1178 );
1179 let Value::StringArray(labels) = &values[0] else {
1180 panic!("labels");
1181 };
1182 assert_eq!(labels.data, vec!["low", "high"]);
1183 let Value::Cell(coeff) = &values[4] else {
1184 panic!("coeff");
1185 };
1186 assert!(
1187 matches!(&coeff.data[1], Value::Struct(st) if st.fields.contains_key("quadratic"))
1188 );
1189 });
1190 }
1191
1192 #[test]
1193 fn classify_empirical_prior_and_missing_labels() {
1194 with_outputs(3, || {
1195 let values = classify(
1196 tensor(vec![0.1, 10.0], vec![2, 1]),
1197 tensor(vec![0.0, 0.2, 10.0, 10.2, 99.0], vec![5, 1]),
1198 tensor(vec![1.0, 1.0, 2.0, 2.0, f64::NAN], vec![5, 1]),
1199 vec![
1200 Value::String("diagLinear".into()),
1201 Value::String("empirical".into()),
1202 ],
1203 );
1204 let Value::Tensor(labels) = &values[0] else {
1205 panic!("labels");
1206 };
1207 assert_eq!(labels.data, vec![1.0, 2.0]);
1208 let Value::Tensor(posterior) = &values[2] else {
1209 panic!("posterior");
1210 };
1211 assert_eq!(posterior.shape, vec![2, 2]);
1212 });
1213 }
1214
1215 #[test]
1216 fn classify_preserves_char_matrix_labels() {
1217 with_outputs(1, || {
1218 let groups = Value::CharArray(
1219 CharArray::new(vec!['a', 'a', 'a', 'a', 'b', 'b', 'b', 'b'], 4, 2).unwrap(),
1220 );
1221 let values = classify(
1222 tensor(vec![0.1, 3.2], vec![2, 1]),
1223 tensor(vec![0.0, 0.4, 3.0, 3.4], vec![4, 1]),
1224 groups,
1225 Vec::new(),
1226 );
1227 let Value::CharArray(labels) = &values[0] else {
1228 panic!("labels");
1229 };
1230 assert_eq!(labels.rows, 2);
1231 assert_eq!(labels.cols, 2);
1232 assert_eq!(labels.data, vec!['a', 'a', 'b', 'b']);
1233 });
1234 }
1235
1236 #[test]
1237 fn classify_accepts_cellstr_labels_and_struct_priors() {
1238 with_outputs(2, || {
1239 let groups = Value::Cell(
1240 CellArray::new(
1241 vec![
1242 Value::CharArray(CharArray::new_row("left")),
1243 Value::CharArray(CharArray::new_row("left")),
1244 Value::CharArray(CharArray::new_row("right")),
1245 Value::CharArray(CharArray::new_row("right")),
1246 ],
1247 4,
1248 1,
1249 )
1250 .unwrap(),
1251 );
1252 let mut prior = StructValue::new();
1253 prior.insert(
1254 "group",
1255 Value::StringArray(
1256 StringArray::new(vec!["left".into(), "right".into()], vec![2, 1]).unwrap(),
1257 ),
1258 );
1259 prior.insert(
1260 "prob",
1261 Value::Tensor(Tensor::new(vec![0.25, 0.75], vec![2, 1]).unwrap()),
1262 );
1263 let values = classify(
1264 tensor(vec![0.1, 3.2], vec![2, 1]),
1265 tensor(vec![0.0, 0.4, 3.0, 3.4], vec![4, 1]),
1266 groups,
1267 vec![Value::String("linear".into()), Value::Struct(prior)],
1268 );
1269 let Value::Cell(labels) = &values[0] else {
1270 panic!("labels");
1271 };
1272 assert_eq!(labels.rows, 2);
1273 assert!(matches!(values[1], Value::Num(err) if err <= EPS));
1274 });
1275 }
1276
1277 #[test]
1278 fn classify_rejects_bad_dimensions() {
1279 let err = block_on(classify_builtin(
1280 tensor(vec![1.0, 2.0], vec![1, 2]),
1281 tensor(vec![1.0, 2.0], vec![2, 1]),
1282 tensor(vec![1.0, 2.0], vec![2, 1]),
1283 Vec::new(),
1284 ))
1285 .unwrap_err();
1286 assert!(err.message.contains("same number of columns"));
1287 }
1288}