Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
39 changes: 37 additions & 2 deletions src/neighbors/knn_classifier.rs
Original file line number Diff line number Diff line change
Expand Up @@ -241,9 +241,9 @@ impl<TX: Number, TY: Number + Ord, X: Array2<TX>, Y: Array1<TY>, D: Distance<Vec
)));
}

if parameters.k <= 1 {
if parameters.k < 1 {
return Err(Failed::fit(&format!(
"k should be > 1, k=[{}]",
"k should be > 0, k=[{}]",
parameters.k
)));
}
Expand Down Expand Up @@ -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() {
Expand Down
Loading