Skip to content
Bill Liao
Go back

K Closest Points to Origin

Edit page

Given an array of points where points[i] = [xi, yi] represents a point on the X-Y plane and an integer k, return the k closest points to the origin (0, 0).

The distance between two points on the X-Y plane is the Euclidean distance (i.e., √(x1 - x2)2 + (y1 - y2)2).

You may return the answer in any order. The answer is guaranteed to be unique (except for the order that it is in).

Example 1:

Input: points = [[1,3],[-2,2]], k = 1 Output: [[-2,2]] Explanation: The distance between (1, 3) and the origin is sqrt(10). The distance between (-2, 2) and the origin is sqrt(8). Since sqrt(8) < sqrt(10), (-2, 2) is closer to the origin. We only want the closest k = 1 points from the origin, so the answer is just [[-2,2]].

Example 2:

Input: points = [[3,3],[5,-1],[-2,4]], k = 2 Output: [[3,3],[-2,4]] Explanation: The answer [[-2,4],[3,3]] would also be accepted.

Constraints:

Approach: Max-Heap of Size k

Algorithm

  1. The squared distance x^2 + y^2 of a point is enough for comparison, avoiding sqrt
  2. Maintain a max-heap of size k, ordered by squared distance descending
  3. For each point, offer it to the heap; whenever the heap exceeds k elements, poll the farthest (largest) one
  4. After processing all points, the heap contains exactly the k closest points
  5. Drain the heap into the result array

Time & Space Complexity

Java Implementation

import java.util.PriorityQueue;

public class KClosestPointsToOrigin {

    public static int[][] kClosest(int[][] points, int k) {
        PriorityQueue<int[]> maxHeap = new PriorityQueue<>(
                (a, b) -> distanceSquared(b) - distanceSquared(a));

        for (int[] p : points) {
            maxHeap.offer(p);
            if (maxHeap.size() > k) {
                maxHeap.poll();
            }
        }

        int[][] result = new int[k][2];
        for (int i = 0; i < k; i++) {
            result[i] = maxHeap.poll();
        }
        return result;
    }

    private static int distanceSquared(int[] p) {
        return p[0] * p[0] + p[1] * p[1];
    }

    // Test method
    public static void main(String[] args) {
        int[][] points1 = {{1, 3}, {-2, 2}};
        int[][] r1 = kClosest(points1, 1);
        System.out.println("kClosest([[1,3],[-2,2]], 1): " + java.util.Arrays.deepToString(r1)); // Expected: [[-2, 2]]

        int[][] points2 = {{3, 3}, {5, -1}, {-2, 4}};
        int[][] r2 = kClosest(points2, 2);
        System.out.println("kClosest([[3,3],[5,-1],[-2,4]], 2): " + java.util.Arrays.deepToString(r2)); // Expected: [[3, 3], [-2, 4]]
    }
}

Alternative Approach: Quickselect

For a faster average case, partition the array around a random pivot by squared distance until the pivot lands at index k - 1:

import java.util.concurrent.ThreadLocalRandom;

public class KClosestPointsToOriginQuickSelect {

    public int[][] kClosest(int[][] points, int k) {
        quickSelect(points, 0, points.length - 1, k);
        int[][] result = new int[k][2];
        System.arraycopy(points, 0, result, 0, k);
        return result;
    }

    private void quickSelect(int[][] points, int lo, int hi, int k) {
        if (lo >= hi) {
            return;
        }
        int pivot = ThreadLocalRandom.current().nextInt(lo, hi + 1);
        int pivotDist = distanceSquared(points[pivot]);
        swap(points, pivot, hi);

        int store = lo;
        for (int i = lo; i < hi; i++) {
            if (distanceSquared(points[i]) < pivotDist) {
                swap(points, store++, i);
            }
        }
        swap(points, store, hi);

        if (store == k - 1) {
            return;
        } else if (store < k - 1) {
            quickSelect(points, store + 1, hi, k);
        } else {
            quickSelect(points, lo, store - 1, k);
        }
    }

    private int distanceSquared(int[] p) {
        return p[0] * p[0] + p[1] * p[1];
    }

    private void swap(int[][] points, int i, int j) {
        int[] temp = points[i];
        points[i] = points[j];
        points[j] = temp;
    }
}

Example Walkthrough

For points = [[3,3],[5,-1],[-2,4]], k = 2:

  1. Squared distances: [18, 26, 20]
  2. Heap starts empty. Add [3,3] (18). Add [5,-1] (26) -> heap [18, 26]
  3. Add [-2,4] (20) -> heap exceeds 2, poll [5,-1] (26). Heap is now [18, 20]
  4. Result: [[3,3],[-2,4]]

Key Points

  1. Max-Heap Trick: Keeping the k closest requires discarding the farthest, so a max-heap at the top is correct
  2. Avoid sqrt: Comparing squared distances is monotonic with actual distance
  3. O(n log k): This beats full sorting at O(n log n)
  4. Answer Uniqueness: The guarantee of unique distances means ties cannot occur in the ordering
  5. Quickselect Alternative: When the heap space O(k) matters less than time, quickselect reaches O(n) on average

Edit page
Share this post:

Previous Post
Number of Boomerangs
Next Post
Valid Square