From 31278918d0a1ba40662e40ae507b0849ce440389 Mon Sep 17 00:00:00 2001 From: poorva0405 Date: Thu, 6 Aug 2026 19:22:57 +0530 Subject: [PATCH] Add K-Nearest Neighbors classifier --- .../machinelearning/KNearestNeighbors.java | 206 ++++++++++++++++ .../KNearestNeighborsTest.java | 220 ++++++++++++++++++ 2 files changed, 426 insertions(+) create mode 100644 src/main/java/com/thealgorithms/machinelearning/KNearestNeighbors.java create mode 100644 src/test/java/com/thealgorithms/machinelearning/KNearestNeighborsTest.java diff --git a/src/main/java/com/thealgorithms/machinelearning/KNearestNeighbors.java b/src/main/java/com/thealgorithms/machinelearning/KNearestNeighbors.java new file mode 100644 index 000000000000..8b2355a00308 --- /dev/null +++ b/src/main/java/com/thealgorithms/machinelearning/KNearestNeighbors.java @@ -0,0 +1,206 @@ +package com.thealgorithms.machinelearning; + +import java.util.ArrayList; +import java.util.Comparator; +import java.util.HashMap; +import java.util.List; +import java.util.Map; + +/** + * K-Nearest Neighbors (KNN) classifier. + * + *

K-Nearest Neighbors is a supervised machine learning algorithm that + * classifies a sample based on the majority class among its {@code k} + * nearest training samples using the Euclidean distance metric. + * + *

The classifier stores the training dataset during the fitting phase and + * predicts class labels for new samples without building an explicit model. + * + * @see + * K-Nearest Neighbors + */ +public final class KNearestNeighbors { + private final int k; + + /** + * Constructs a K-Nearest Neighbors classifier with the specified number + * of neighbors. + * + * @param k the number of nearest neighbors to consider during prediction + */ + public KNearestNeighbors(int k) { + + if (k <= 0) { + throw new IllegalArgumentException("k must be greater than 0."); + } + + this.k = k; + } + + /** + * Represents a neighboring training sample and its distance from the test sample. + */ + private static final class Neighbor { + + private final double distance; + + private final int label; + + Neighbor(double distance, int label) { + this.distance = distance; + this.label = label; + } + } + + private double[][] trainingFeatures; + private int[] trainingLabels; + private int numFeatures; + + /** + * Fits the classifier using the provided training dataset. + * + *

The training feature vectors and their corresponding class labels are + * stored for use during prediction. + * + * @param features the training feature vectors + * @param labels the corresponding class labels + */ + public void fit(double[][] features, int[] labels) { + + if (features == null || labels == null) { + throw new IllegalArgumentException("Features and labels cannot be null."); + } + + if (features.length == 0 || labels.length == 0) { + throw new IllegalArgumentException("Features and labels cannot be empty."); + } + + if (features.length != labels.length) { + throw new IllegalArgumentException("Features and labels must have the same length."); + } + + if (features[0] == null) { + throw new IllegalArgumentException("Feature vectors cannot be null."); + } + + numFeatures = features[0].length; + + if (numFeatures == 0) { + throw new IllegalArgumentException("Feature vectors cannot be empty."); + } + + for (double[] sample : features) { + if (sample == null) { + throw new IllegalArgumentException("Feature vectors cannot be null."); + } + + if (sample.length != numFeatures) { + throw new IllegalArgumentException("All feature vectors must have the same dimension."); + } + } + + this.trainingFeatures = features; + this.trainingLabels = labels; + } + + /** + * Computes the Euclidean distance between two feature vectors. + * + * @param first the first feature vector + * @param second the second feature vector + * @return the Euclidean distance between the two vectors + */ + private static double euclideanDistance(double[] first, double[] second) { + double sum = 0.0; + + for (int i = 0; i < first.length; i++) { + double difference = first[i] - second[i]; + sum += difference * difference; + } + + return Math.sqrt(sum); + } + + /** + * Predicts the class label for a single sample. + * + *

The prediction is made by finding the {@code k} nearest neighbors + * among the training samples and selecting the class with the highest + * number of votes. In the event of a tie, the smaller class label is + * returned. + * + * @param testPoint the sample to classify + * @return the predicted class label + */ + public int predict(double[] testPoint) { + if (trainingFeatures == null || trainingLabels == null) { + throw new IllegalStateException("Classifier has not been fitted."); + } + + if(testPoint == null) { + throw new IllegalArgumentException("Sample cannot be null."); + } + + if (testPoint.length != numFeatures) { + throw new IllegalArgumentException("Sample length must match training feature count."); + } + + List neighbors = new ArrayList<>(trainingFeatures.length); + + for (int i = 0; i < trainingFeatures.length; i++) { + double distance = euclideanDistance(trainingFeatures[i], testPoint); + neighbors.add(new Neighbor(distance, trainingLabels[i])); + } + + + neighbors.sort(Comparator.comparingDouble(neighbor -> neighbor.distance)); + + Map votes = new HashMap<>(); + + if (k > trainingFeatures.length) { + throw new IllegalArgumentException("k cannot be greater than the number of training samples."); + } + + for (int i = 0; i < k; i++) { + int label = neighbors.get(i).label; + votes.merge(label, 1, Integer::sum); + } + + int predictedLabel = -1; + int maxVotes = -1; + + for (Map.Entry entry : votes.entrySet()) { + int label = entry.getKey(); + int count = entry.getValue(); + + if (count > maxVotes || (count == maxVotes && label < predictedLabel)) { + maxVotes = count; + predictedLabel = label; + } + } + + return predictedLabel; + } + + + /** + * Predicts class labels for multiple samples. + * + * @param samples the samples to classify + * @return an array containing the predicted class label for each sample + */ + public int[] predict(double[][] samples) { + + if (samples == null) { + throw new IllegalArgumentException("Samples cannot be null."); + } + + int[] predictions = new int[samples.length]; + + for (int i = 0; i < samples.length; i++) { + predictions[i] = predict(samples[i]); + } + + return predictions; + } +} diff --git a/src/test/java/com/thealgorithms/machinelearning/KNearestNeighborsTest.java b/src/test/java/com/thealgorithms/machinelearning/KNearestNeighborsTest.java new file mode 100644 index 000000000000..503a15d45b1e --- /dev/null +++ b/src/test/java/com/thealgorithms/machinelearning/KNearestNeighborsTest.java @@ -0,0 +1,220 @@ +package com.thealgorithms.machinelearning; + +import static org.junit.jupiter.api.Assertions.assertArrayEquals; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; + +import org.junit.jupiter.api.Test; + +public class KNearestNeighborsTest { + + @Test + void predictsCorrectClassOnSeparableDataset() { + double[][] features = { + {0, 0}, + {1, 1}, + {8, 8}, + {9, 9}, + }; + + int[] labels = {0, 0, 1, 1}; + + KNearestNeighbors knn = new KNearestNeighbors(3); + knn.fit(features, labels); + + assertEquals(0, knn.predict(new double[] {1.5, 1.5})); + assertEquals(1, knn.predict(new double[] {8.5, 8.5})); + } + + @Test + void predictsBatchClassesOnSeparableDataset() { + double[][] features = { + {0, 0}, + {1, 1}, + {8, 8}, + {9, 9}, + }; + + int[] labels = {0, 0, 1, 1}; + + KNearestNeighbors knn = new KNearestNeighbors(3); + knn.fit(features, labels); + + double[][] samples = {{1.5, 1.5}, {8.5, 8.5}}; + int[] predictions = knn.predict(samples); + + assertArrayEquals(new int[] {0, 1}, predictions); + } + + @Test + void throwsExceptionWhenKIsNotPositive() { + assertThrows(IllegalArgumentException.class, () -> new KNearestNeighbors(-4)); + assertThrows(IllegalArgumentException.class, () -> new KNearestNeighbors(0)); + } + + @Test + void throwsExceptionWhenKIsGreaterThanNumberOfTrainingSamples() { + double[][] features = { + {0, 0}, + {1, 1}, + {8, 8}, + {9, 9}, + }; + + int[] labels = {0, 0, 1, 1}; + KNearestNeighbors knn = new KNearestNeighbors(7); + knn.fit(features, labels); + + assertThrows(IllegalArgumentException.class, () -> knn.predict(new double[] {1.5, 1.5})); + } + + @Test + void nullFeaturesArrayThrowsIllegalArgumentException() { + KNearestNeighbors knn = new KNearestNeighbors(1); + + assertThrows(IllegalArgumentException.class, () -> knn.fit(null, new int[] {0, 1})); + } + + @Test + void emptyFeaturesArrayThrowsIllegalArgumentException() { + KNearestNeighbors knn = new KNearestNeighbors(1); + + assertThrows(IllegalArgumentException.class, () -> knn.fit(new double[][] {}, new int[] {})); + } + + @Test + void nullLabelsArrayThrowsIllegalArgumentException() { + KNearestNeighbors knn = new KNearestNeighbors(1); + + assertThrows(IllegalArgumentException.class, () -> knn.fit(new double[][] {{1, 1}, {2, 2}}, null)); + } + + @Test + void emptyLabelsArrayThrowsIllegalArgumentException() { + KNearestNeighbors knn = new KNearestNeighbors(1); + + assertThrows(IllegalArgumentException.class, () -> knn.fit(new double[][] {{1, 1}, {2, 2}}, new int[] {})); + } + + @Test + void mismatchedFeatureAndLabelLengthsThrowsIllegalArgumentException() { + KNearestNeighbors knn = new KNearestNeighbors(3); + + assertThrows(IllegalArgumentException.class, () -> knn.fit(new double[][] {{0, 0}, {1, 1}, {8, 8}, {9, 9}}, new int[] {0, 0, 1})); + } + + @Test + void emptyFeatureVectorThrowsIllegalArgumentException() { + KNearestNeighbors knn = new KNearestNeighbors(1); + + assertThrows(IllegalArgumentException.class,() -> knn.fit(new double[][] {{}}, new int[] {0})); + } + + @Test + void nullFeatureSampleThrowsIllegalArgumentException() { + KNearestNeighbors knn = new KNearestNeighbors(2); + + assertThrows(IllegalArgumentException.class, () -> knn.fit(new double[][] {{0, 0}, null, {8, 8}}, new int[] {0, 0, 1})); + } + + @Test + void mismatchedDimensionThrowsIllegalArgumentException() { + KNearestNeighbors knn = new KNearestNeighbors(2); + + assertThrows(IllegalArgumentException.class, () -> knn.fit(new double[][] {{0, 0}, {1, 1, 2}, {8, 8}}, new int[] {0, 0, 1})); + } + + @Test + void predictBeforeFitThrowsIllegalStateException() { + KNearestNeighbors knn = new KNearestNeighbors(3); + + assertThrows(IllegalStateException.class, () -> knn.predict(new double[] {1, 1})); + } + + @Test + void nullTestPointThrowsIllegalArgumentException() { + double[][] features = { + {0, 0}, + {1, 1}, + {8, 8}, + {9, 9}, + }; + + int[] labels = {0, 0, 1, 1}; + + KNearestNeighbors knn = new KNearestNeighbors(3); + knn.fit(features, labels); + + assertThrows(IllegalArgumentException.class, () -> knn.predict((double[]) null)); + assertThrows(IllegalArgumentException.class, () -> knn.predict((double[][]) null)); + } + + @Test + void mismatchedTestPointLengthThrowsIllegalArgumentException() { + double[][] features = { + {0, 0}, + {1, 1}, + {8, 8}, + {9, 9}, + }; + + int[] labels = {0, 0, 1, 1}; + + KNearestNeighbors knn = new KNearestNeighbors(3); + knn.fit(features, labels); + + assertThrows(IllegalArgumentException.class, () -> knn.predict(new double[] {1.5, 1.5, 1.5})); + } + + @Test + void tieBreakReturnsSmallerLabel() { + double[][] features = { + {0, 0}, + {0, 2}, + {2, 0}, + {2, 2} + }; + + int[] labels = {0, 1, 0, 1}; + + KNearestNeighbors knn = new KNearestNeighbors(4); + knn.fit(features, labels); + + assertEquals(0, knn.predict(new double[] {1, 1})); + } + + @Test + void nullBatchSampleThrowsIllegalArgumentException() { + double[][] features = { + {0, 0}, + {1, 1}, + {8, 8}, + {9, 9}, + }; + + int[] labels = {0, 0, 1, 1}; + + KNearestNeighbors knn = new KNearestNeighbors(3); + knn.fit(features, labels); + + assertThrows(IllegalArgumentException.class, () -> knn.predict(new double[][] {{1.5, 1.5}, null})); + } + + @Test + void invalidBatchSampleThrowsIllegalArgumentException() { + double[][] features = { + {0, 0}, + {1, 1}, + {8, 8}, + {9, 9}, + }; + + int[] labels = {0, 0, 1, 1}; + + KNearestNeighbors knn = new KNearestNeighbors(3); + knn.fit(features, labels); + + assertThrows(IllegalArgumentException.class, () -> knn.predict(new double[][] {{1.5, 1.5}, {8.5, 8.5, 8.5}})); + } + +} \ No newline at end of file