1use 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#[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#[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#[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#[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#[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#[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}