Skip to main content

sklears_python/metrics/
classification.rs

1//! Python bindings for classification metrics
2
3use super::common::*;
4use numpy::{PyArray2, PyReadonlyArray1};
5use scirs2_core::ndarray::Array1;
6use sklears_metrics::basic_metrics::{
7    accuracy_score as skl_accuracy, confusion_matrix as skl_confusion_matrix, f1_score as skl_f1,
8    precision_score as skl_precision, recall_score as skl_recall,
9};
10use std::collections::HashMap;
11
12/// Calculate accuracy score for classification
13#[pyfunction]
14#[pyo3(signature = (y_true, y_pred, normalize=true, sample_weight=None))]
15pub fn accuracy_score(
16    y_true: PyReadonlyArray1<i32>,
17    y_pred: PyReadonlyArray1<i32>,
18    normalize: bool,
19    sample_weight: Option<PyReadonlyArray1<f64>>,
20) -> PyResult<f64> {
21    let _ = sample_weight;
22    let yt = Array1::from_vec(y_true.as_array().to_vec());
23    let yp = Array1::from_vec(y_pred.as_array().to_vec());
24
25    let yt_slice = yt
26        .as_slice()
27        .ok_or_else(|| PyValueError::new_err("internal error: array not contiguous"))?;
28    let yp_slice = yp
29        .as_slice()
30        .ok_or_else(|| PyValueError::new_err("internal error: array not contiguous"))?;
31    validate_int_arrays_same_length(yt_slice, yp_slice)?;
32
33    match skl_accuracy(&yt, &yp) {
34        Ok(acc) => Ok(if normalize {
35            acc
36        } else {
37            acc * yt.len() as f64
38        }),
39        Err(e) => Err(PyValueError::new_err(format!("accuracy_score: {}", e))),
40    }
41}
42
43/// Calculate precision score for binary classification
44#[pyfunction]
45#[pyo3(signature = (y_true, y_pred, labels=None, pos_label=1, average="binary", sample_weight=None, zero_division="warn"))]
46pub fn precision_score(
47    y_true: PyReadonlyArray1<i32>,
48    y_pred: PyReadonlyArray1<i32>,
49    labels: Option<PyReadonlyArray1<i32>>,
50    pos_label: i32,
51    average: &str,
52    sample_weight: Option<PyReadonlyArray1<f64>>,
53    zero_division: &str,
54) -> PyResult<f64> {
55    let _ = (labels, average, sample_weight, zero_division);
56    let yt = Array1::from_vec(y_true.as_array().to_vec());
57    let yp = Array1::from_vec(y_pred.as_array().to_vec());
58
59    let yt_slice = yt
60        .as_slice()
61        .ok_or_else(|| PyValueError::new_err("internal error: array not contiguous"))?;
62    let yp_slice = yp
63        .as_slice()
64        .ok_or_else(|| PyValueError::new_err("internal error: array not contiguous"))?;
65    validate_int_arrays_same_length(yt_slice, yp_slice)?;
66
67    match skl_precision(&yt, &yp, Some(pos_label)) {
68        Ok(v) => Ok(v),
69        Err(e) => Err(PyValueError::new_err(format!("precision_score: {}", e))),
70    }
71}
72
73/// Calculate recall score for binary classification
74#[pyfunction]
75#[pyo3(signature = (y_true, y_pred, labels=None, pos_label=1, average="binary", sample_weight=None, zero_division="warn"))]
76pub fn recall_score(
77    y_true: PyReadonlyArray1<i32>,
78    y_pred: PyReadonlyArray1<i32>,
79    labels: Option<PyReadonlyArray1<i32>>,
80    pos_label: i32,
81    average: &str,
82    sample_weight: Option<PyReadonlyArray1<f64>>,
83    zero_division: &str,
84) -> PyResult<f64> {
85    let _ = (labels, average, sample_weight, zero_division);
86    let yt = Array1::from_vec(y_true.as_array().to_vec());
87    let yp = Array1::from_vec(y_pred.as_array().to_vec());
88
89    let yt_slice = yt
90        .as_slice()
91        .ok_or_else(|| PyValueError::new_err("internal error: array not contiguous"))?;
92    let yp_slice = yp
93        .as_slice()
94        .ok_or_else(|| PyValueError::new_err("internal error: array not contiguous"))?;
95    validate_int_arrays_same_length(yt_slice, yp_slice)?;
96
97    match skl_recall(&yt, &yp, Some(pos_label)) {
98        Ok(v) => Ok(v),
99        Err(e) => Err(PyValueError::new_err(format!("recall_score: {}", e))),
100    }
101}
102
103/// Calculate F1 score for binary classification
104#[pyfunction]
105#[pyo3(signature = (y_true, y_pred, labels=None, pos_label=1, average="binary", sample_weight=None, zero_division="warn"))]
106pub fn f1_score(
107    y_true: PyReadonlyArray1<i32>,
108    y_pred: PyReadonlyArray1<i32>,
109    labels: Option<PyReadonlyArray1<i32>>,
110    pos_label: i32,
111    average: &str,
112    sample_weight: Option<PyReadonlyArray1<f64>>,
113    zero_division: &str,
114) -> PyResult<f64> {
115    let _ = (labels, average, sample_weight, zero_division);
116    let yt = Array1::from_vec(y_true.as_array().to_vec());
117    let yp = Array1::from_vec(y_pred.as_array().to_vec());
118
119    let yt_slice = yt
120        .as_slice()
121        .ok_or_else(|| PyValueError::new_err("internal error: array not contiguous"))?;
122    let yp_slice = yp
123        .as_slice()
124        .ok_or_else(|| PyValueError::new_err("internal error: array not contiguous"))?;
125    validate_int_arrays_same_length(yt_slice, yp_slice)?;
126
127    match skl_f1(&yt, &yp, Some(pos_label)) {
128        Ok(v) => Ok(v),
129        Err(e) => Err(PyValueError::new_err(format!("f1_score: {}", e))),
130    }
131}
132
133/// Calculate confusion matrix
134#[pyfunction]
135#[pyo3(signature = (y_true, y_pred, labels=None, sample_weight=None, normalize=None))]
136pub fn confusion_matrix(
137    py: Python,
138    y_true: PyReadonlyArray1<i32>,
139    y_pred: PyReadonlyArray1<i32>,
140    labels: Option<PyReadonlyArray1<i32>>,
141    sample_weight: Option<PyReadonlyArray1<f64>>,
142    normalize: Option<&str>,
143) -> PyResult<Py<PyArray2<i64>>> {
144    let _ = (labels, sample_weight, normalize);
145    let yt = Array1::from_vec(y_true.as_array().to_vec());
146    let yp = Array1::from_vec(y_pred.as_array().to_vec());
147
148    let yt_slice = yt
149        .as_slice()
150        .ok_or_else(|| PyValueError::new_err("internal error: array not contiguous"))?;
151    let yp_slice = yp
152        .as_slice()
153        .ok_or_else(|| PyValueError::new_err("internal error: array not contiguous"))?;
154    validate_int_arrays_same_length(yt_slice, yp_slice)?;
155
156    match skl_confusion_matrix(&yt, &yp) {
157        Ok(cm) => {
158            let cm_i64 = cm.mapv(|v| v as i64);
159            Ok(PyArray2::from_array(py, &cm_i64).unbind())
160        }
161        Err(e) => Err(PyValueError::new_err(format!("confusion_matrix: {}", e))),
162    }
163}
164
165/// Calculate classification report (returns a nested dict when output_dict=True).
166/// Note: `digits` and `zero_division` are accepted for API compatibility but ignored.
167#[pyfunction]
168#[pyo3(signature = (y_true, y_pred, labels=None, target_names=None, sample_weight=None, output_dict=true))]
169pub fn classification_report(
170    y_true: PyReadonlyArray1<i32>,
171    y_pred: PyReadonlyArray1<i32>,
172    labels: Option<PyReadonlyArray1<i32>>,
173    target_names: Option<Vec<String>>,
174    sample_weight: Option<PyReadonlyArray1<f64>>,
175    output_dict: bool,
176) -> PyResult<HashMap<String, HashMap<String, f64>>> {
177    let _ = (labels, target_names, sample_weight);
178    if !output_dict {
179        return Err(PyValueError::new_err(
180            "String output not supported; use output_dict=True.",
181        ));
182    }
183
184    let yt = Array1::from_vec(y_true.as_array().to_vec());
185    let yp = Array1::from_vec(y_pred.as_array().to_vec());
186
187    let yt_slice = yt
188        .as_slice()
189        .ok_or_else(|| PyValueError::new_err("internal error: array not contiguous"))?;
190    let yp_slice = yp
191        .as_slice()
192        .ok_or_else(|| PyValueError::new_err("internal error: array not contiguous"))?;
193    validate_int_arrays_same_length(yt_slice, yp_slice)?;
194
195    let pos_label = *yt.iter().max().unwrap_or(&1);
196    let support = yt.len() as f64;
197
198    let precision = skl_precision(&yt, &yp, Some(pos_label)).unwrap_or(0.0);
199    let recall = skl_recall(&yt, &yp, Some(pos_label)).unwrap_or(0.0);
200    let f1 = skl_f1(&yt, &yp, Some(pos_label)).unwrap_or(0.0);
201
202    let mut class_entry = HashMap::new();
203    class_entry.insert("precision".to_string(), precision);
204    class_entry.insert("recall".to_string(), recall);
205    class_entry.insert("f1-score".to_string(), f1);
206    class_entry.insert("support".to_string(), support);
207
208    let mut report = HashMap::new();
209    report.insert(pos_label.to_string(), class_entry);
210    Ok(report)
211}