diff --git a/CHANGELOG.md b/CHANGELOG.md index c250a5db..8651734f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,9 +7,11 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [0.6.16] ### Fixed - `predict`, `predict_proba`, `predict_oob`, `predict_matrix` and `decision_function` now return `Err(Failed)` instead of panicking when `fit` has not run (#469, #470). Models return the same unfitted error even after deserialization, when the state fields are `None`. The guarded methods cover the decision trees, the random forests, the extra-trees regressor, the KNN classifier and regressor, `SVC`, `MultiClassSVC`, `SVR`, the linear models, the naive Bayes classifiers and `XGRegressor`. +- `KNNClassifier::fit` now accepts `k = 1`, as `KNNRegressor` already did (#476). `k = 0` is still rejected. ### Changed - The error text for an unfitted `MultiClassSVC::predict` changed from "MultiClassSVC is not fitted" to the common "'fit' should be called before calling 'predict'". Code that matched the old string must be updated. +- The error text for an invalid `k` in `KNNClassifier::fit` changed from "k should be > 1" to "k should be > 0", the same text `KNNRegressor` uses. Code that matched the old string must be updated. ## [0.6.15] ### Fixed diff --git a/src/neighbors/knn_classifier.rs b/src/neighbors/knn_classifier.rs index b1c57e14..fbfc306e 100644 --- a/src/neighbors/knn_classifier.rs +++ b/src/neighbors/knn_classifier.rs @@ -241,9 +241,9 @@ impl, Y: Array1, D: Distance 1, k=[{}]", + "k should be > 0, k=[{}]", parameters.k ))); } @@ -421,6 +421,41 @@ mod tests { assert_eq!(vec![3], y_hat); } + #[test] + fn knn_fit_predict_k1() { + let x = DenseMatrix::from_2d_array(&[&[1.], &[2.], &[3.], &[4.], &[5.]]).unwrap(); + let y = vec![2, 3, 2, 3, 2]; + let one_hot = vec![ + vec![1., 0.], + vec![0., 1.], + vec![1., 0.], + vec![0., 1.], + vec![1., 0.], + ]; + + for algorithm in [KNNAlgorithmName::CoverTree, KNNAlgorithmName::LinearSearch] { + let knn = KNNClassifier::fit( + &x, + &y, + KNNClassifierParameters::default() + .with_k(1) + .with_algorithm(algorithm), + ) + .unwrap(); + + assert_eq!(y, knn.predict(&x).unwrap()); + assert_eq!(one_hot, knn.predict_proba(&x).unwrap()); + } + } + + #[test] + fn knn_fit_k0_fails() { + let x = DenseMatrix::from_2d_array(&[&[1.], &[2.]]).unwrap(); + let y = vec![2, 3]; + + assert!(KNNClassifier::fit(&x, &y, KNNClassifierParameters::default().with_k(0)).is_err()); + } + // New 8 tests (2026-03-19) #[test] fn knn_predict_proba_valid() {