Skip to main content

runmat_runtime/builtins/stats/ml/
classification.rs

1//! Discriminant classification compatibility surface.
2
3use 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}