From 365fdb6c83119362a3ed56d2c3c3a3eafeb72b1e Mon Sep 17 00:00:00 2001 From: Rajat Singh Date: Thu, 6 Aug 2026 21:02:11 +0530 Subject: [PATCH 1/6] Feat: Add K-Nearest Neighbors (KNN) classification algorithm (#7562) Implements the K-Nearest Neighbors (KNN) classification algorithm from scratch using native Java utility structures. Key Features: - Computes spatial vector variance using the Euclidean Distance metric calculation. - Leverages a custom internal data layout to manage feature/label pairs cleanly. - Processes classification matching via a frequency-based majority vote resolution. - Includes a standalone execution test routine within the class main entry channel. --- .../machinelearning/machinelearning/KNN.java | 74 +++++++++++++++++++ 1 file changed, 74 insertions(+) create mode 100644 src/main/java/com/thealgorithms/machinelearning/machinelearning/KNN.java diff --git a/src/main/java/com/thealgorithms/machinelearning/machinelearning/KNN.java b/src/main/java/com/thealgorithms/machinelearning/machinelearning/KNN.java new file mode 100644 index 000000000000..c7d095916bfe --- /dev/null +++ b/src/main/java/com/thealgorithms/machinelearning/machinelearning/KNN.java @@ -0,0 +1,74 @@ +package com.thealgorithms.machinelearning; + +import java.util.ArrayList; +import java.util.Collections; +import java.util.Comparator; +import java.util.List; + +public class KNN { + public static class DataPoint { + double[] features; + String label; + + public DataPoint(double[] features, String label) { + this.features = features; + this.label = label; + } + } + + private static class DistancePair { + double distance; + String label; + + public DistancePair(double distance, String label) { + this.distance = distance; + this.label = label; + } + } + + public static double calculateEuclideanDistance(double[] point1, double[] point2) { + double sum = 0.0; + for (int i = 0; i < point1.length; i++) { + sum += Math.pow(point1[i] - point2[i], 2); + } + return Math.sqrt(sum); + } + + public static String classify(List dataset, double[] queryPoint, int k) { + List distances = new ArrayList<>(); + + for (DataPoint p : dataset) { + double dist = calculateEuclideanDistance(p.features, queryPoint); + distances.add(new DistancePair(dist, p.label)); + } + + distances.sort(Comparator.comparingDouble(d -> d.distance)); + + List topKLabels = new ArrayList<>(); + for (int i = 0; i < Math.min(k, distances.size()); i++) { + topKLabels.add(distances.get(i).label); + } + + String bestLabel = null; + int maxCount = -1; + for (String label : topKLabels) { + int count = Collections.frequency(topKLabels, label); + if (count > maxCount) { + maxCount = count; + bestLabel = label; + } + } + return bestLabel; + } + + public static void main(String[] args) { + List trainData = new ArrayList<>(); + trainData.add(new DataPoint(new double[]{1.0, 2.0}, "ClassA")); + trainData.add(new DataPoint(new double[]{2.0, 3.0}, "ClassA")); + trainData.add(new DataPoint(new double[]{7.0, 8.0}, "ClassB")); + + double[] target = new double[]{1.5, 2.5}; + String prediction = classify(trainData, target, 3); + System.out.println("Predicted Category: " + prediction); + } +} From e87160176d80767952ef2a25d949acae3171a056 Mon Sep 17 00:00:00 2001 From: Rajat Singh Date: Thu, 6 Aug 2026 21:21:55 +0530 Subject: [PATCH 2/6] Delete src/main/java/com/thealgorithms/machinelearning/machinelearning/KNN.java --- .../machinelearning/machinelearning/KNN.java | 74 ------------------- 1 file changed, 74 deletions(-) delete mode 100644 src/main/java/com/thealgorithms/machinelearning/machinelearning/KNN.java diff --git a/src/main/java/com/thealgorithms/machinelearning/machinelearning/KNN.java b/src/main/java/com/thealgorithms/machinelearning/machinelearning/KNN.java deleted file mode 100644 index c7d095916bfe..000000000000 --- a/src/main/java/com/thealgorithms/machinelearning/machinelearning/KNN.java +++ /dev/null @@ -1,74 +0,0 @@ -package com.thealgorithms.machinelearning; - -import java.util.ArrayList; -import java.util.Collections; -import java.util.Comparator; -import java.util.List; - -public class KNN { - public static class DataPoint { - double[] features; - String label; - - public DataPoint(double[] features, String label) { - this.features = features; - this.label = label; - } - } - - private static class DistancePair { - double distance; - String label; - - public DistancePair(double distance, String label) { - this.distance = distance; - this.label = label; - } - } - - public static double calculateEuclideanDistance(double[] point1, double[] point2) { - double sum = 0.0; - for (int i = 0; i < point1.length; i++) { - sum += Math.pow(point1[i] - point2[i], 2); - } - return Math.sqrt(sum); - } - - public static String classify(List dataset, double[] queryPoint, int k) { - List distances = new ArrayList<>(); - - for (DataPoint p : dataset) { - double dist = calculateEuclideanDistance(p.features, queryPoint); - distances.add(new DistancePair(dist, p.label)); - } - - distances.sort(Comparator.comparingDouble(d -> d.distance)); - - List topKLabels = new ArrayList<>(); - for (int i = 0; i < Math.min(k, distances.size()); i++) { - topKLabels.add(distances.get(i).label); - } - - String bestLabel = null; - int maxCount = -1; - for (String label : topKLabels) { - int count = Collections.frequency(topKLabels, label); - if (count > maxCount) { - maxCount = count; - bestLabel = label; - } - } - return bestLabel; - } - - public static void main(String[] args) { - List trainData = new ArrayList<>(); - trainData.add(new DataPoint(new double[]{1.0, 2.0}, "ClassA")); - trainData.add(new DataPoint(new double[]{2.0, 3.0}, "ClassA")); - trainData.add(new DataPoint(new double[]{7.0, 8.0}, "ClassB")); - - double[] target = new double[]{1.5, 2.5}; - String prediction = classify(trainData, target, 3); - System.out.println("Predicted Category: " + prediction); - } -} From e98990b321ef6f0ddec681cee0a32d9eb8ef3bf0 Mon Sep 17 00:00:00 2001 From: Rajat Singh Date: Thu, 6 Aug 2026 21:23:33 +0530 Subject: [PATCH 3/6] Implement K-Nearest Neighbors algorithm This implementation of the K-Nearest Neighbors algorithm includes methods for calculating Euclidean distance and classifying data points based on their nearest neighbors. --- .../thealgorithms/machinelearning/KNN.java | 74 +++++++++++++++++++ 1 file changed, 74 insertions(+) create mode 100644 src/main/java/com/thealgorithms/machinelearning/KNN.java diff --git a/src/main/java/com/thealgorithms/machinelearning/KNN.java b/src/main/java/com/thealgorithms/machinelearning/KNN.java new file mode 100644 index 000000000000..2d56a1821203 --- /dev/null +++ b/src/main/java/com/thealgorithms/machinelearning/KNN.java @@ -0,0 +1,74 @@ +package com.thealgorithms.machinelearning; + +import java.util.ArrayList; +import java.util.Collections; +import java.util.Comparator; +import java.util.List; + +public class KNN { + public static class DataPoint { + double[] features; + String label; + + public DataPoint(double[] features, String label) { + this.features = features; + this.label = label; + } + } + + private static class DistancePair { + double distance; + String label; + + public DistancePair(double distance, String label) { + this.distance = distance; + this.label = label; + } + } + + public static double calculateEuclideanDistance(double[] point1, double[] point2) { + double sum = 0.0; + for (int i = 0; i < point1.length; i++) { + sum += Math.pow(point1[i] - point2[i], 2); + } + return Math.sqrt(sum); + } + + public static String classify(List dataset, double[] queryPoint, int k) { + List distances = new ArrayList<>(); + + for (DataPoint p : dataset) { + double dist = calculateEuclideanDistance(p.features, queryPoint); + distances.add(new DistancePair(dist, p.label)); + } + + distances.sort(Comparator.comparingDouble(d -> d.distance)); + + List topKLabels = new ArrayList<>(); + for (int i = 0; i < Math.min(k, distances.size()); i++) { + topKLabels.add(distances.get(i).label); + } + + String bestLabel = null; + int maxCount = -1; + for (String label : topKLabels) { + int count = Collections.frequency(topKLabels, label); + if (count > maxCount) { + maxCount = count; + bestLabel = label; + } + } + return bestLabel; + } + + public static void main(String[] args) { + List trainData = new ArrayList<>(); + trainData.add(new DataPoint(new double[] {1.0, 2.0}, "ClassA")); + trainData.add(new DataPoint(new double[] {2.0, 3.0}, "ClassA")); + trainData.add(new DataPoint(new double[] {7.0, 8.0}, "ClassB")); + + double[] target = new double[] {1.5, 2.5}; + String prediction = classify(trainData, target, 3); + System.out.println("Predicted Category: " + prediction); + } +} From 75deff406d4e3bdd3903d4a922f4b19781f4c906 Mon Sep 17 00:00:00 2001 From: Rajat Singh Date: Thu, 6 Aug 2026 21:25:22 +0530 Subject: [PATCH 4/6] Add KNNTest class for KNN classification testing --- .../machinelearning/KNNTest.java | 20 +++++++++++++++++++ 1 file changed, 20 insertions(+) create mode 100644 src/main/java/com/thealgorithms/machinelearning/KNNTest.java diff --git a/src/main/java/com/thealgorithms/machinelearning/KNNTest.java b/src/main/java/com/thealgorithms/machinelearning/KNNTest.java new file mode 100644 index 000000000000..d36791361d89 --- /dev/null +++ b/src/main/java/com/thealgorithms/machinelearning/KNNTest.java @@ -0,0 +1,20 @@ +package com.thealgorithms.machinelearning; + +import java.util.ArrayList; +import java.util.List; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +public class KNNTest { + @Test + public void testClassifyStandard() { + List dataset = new ArrayList<>(); + dataset.add(new KNN.DataPoint(new double[] {1.0, 1.0}, "GroupA")); + dataset.add(new KNN.DataPoint(new double[] {1.5, 2.0}, "GroupA")); + dataset.add(new KNN.DataPoint(new double[] {8.0, 9.0}, "GroupB")); + + double[] target = new double[] {1.2, 1.3}; + String result = KNN.classify(dataset, target, 2); + Assertions.assertEquals("GroupA", result); + } +} From 4ca2fb52d7389e0a9428b5e7dd9256aa62488350 Mon Sep 17 00:00:00 2001 From: Rajat Singh Date: Thu, 6 Aug 2026 21:26:24 +0530 Subject: [PATCH 5/6] Delete src/main/java/com/thealgorithms/machinelearning/KNNTest.java --- .../machinelearning/KNNTest.java | 20 ------------------- 1 file changed, 20 deletions(-) delete mode 100644 src/main/java/com/thealgorithms/machinelearning/KNNTest.java diff --git a/src/main/java/com/thealgorithms/machinelearning/KNNTest.java b/src/main/java/com/thealgorithms/machinelearning/KNNTest.java deleted file mode 100644 index d36791361d89..000000000000 --- a/src/main/java/com/thealgorithms/machinelearning/KNNTest.java +++ /dev/null @@ -1,20 +0,0 @@ -package com.thealgorithms.machinelearning; - -import java.util.ArrayList; -import java.util.List; -import org.junit.jupiter.api.Assertions; -import org.junit.jupiter.api.Test; - -public class KNNTest { - @Test - public void testClassifyStandard() { - List dataset = new ArrayList<>(); - dataset.add(new KNN.DataPoint(new double[] {1.0, 1.0}, "GroupA")); - dataset.add(new KNN.DataPoint(new double[] {1.5, 2.0}, "GroupA")); - dataset.add(new KNN.DataPoint(new double[] {8.0, 9.0}, "GroupB")); - - double[] target = new double[] {1.2, 1.3}; - String result = KNN.classify(dataset, target, 2); - Assertions.assertEquals("GroupA", result); - } -} From eace48c6a7a1000dcd3d7d81a45e0759435e540e Mon Sep 17 00:00:00 2001 From: Rajat Singh Date: Thu, 6 Aug 2026 21:28:31 +0530 Subject: [PATCH 6/6] Add unit test for KNN classification --- .../machinelearning/KNNTest.java | 20 +++++++++++++++++++ 1 file changed, 20 insertions(+) create mode 100644 src/test/java/com/thealgorithms/machinelearning/KNNTest.java diff --git a/src/test/java/com/thealgorithms/machinelearning/KNNTest.java b/src/test/java/com/thealgorithms/machinelearning/KNNTest.java new file mode 100644 index 000000000000..d36791361d89 --- /dev/null +++ b/src/test/java/com/thealgorithms/machinelearning/KNNTest.java @@ -0,0 +1,20 @@ +package com.thealgorithms.machinelearning; + +import java.util.ArrayList; +import java.util.List; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +public class KNNTest { + @Test + public void testClassifyStandard() { + List dataset = new ArrayList<>(); + dataset.add(new KNN.DataPoint(new double[] {1.0, 1.0}, "GroupA")); + dataset.add(new KNN.DataPoint(new double[] {1.5, 2.0}, "GroupA")); + dataset.add(new KNN.DataPoint(new double[] {8.0, 9.0}, "GroupB")); + + double[] target = new double[] {1.2, 1.3}; + String result = KNN.classify(dataset, target, 2); + Assertions.assertEquals("GroupA", result); + } +}