DSA interview questionsQuestion 74 of 147
DSA interview question · Question 74 of 147
K Closest Points to Origin: Bounded Max-Heap and Quickselect
Short answer
Compare squared distances x*x + y*y, since square roots do not change the order. Keep a max-heap of the k closest points seen so far (negated distances in heapq); for each new point, push it and pop the farthest when the heap exceeds k. That is O(n log k) time and O(k) space. Quickselect partitions around a pivot distance until the first k positions hold the k closest, which is O(n) on average but O(n^2) in the worst case.
On this page
Problem
You are given a list of points on a plane, each a pair of integers [x, y], and an integer k. Return the k points closest to the origin (0, 0) by ordinary straight-line (Euclidean) distance. The answer may be in any order. You may assume the k-th and (k+1)-th closest points are not tied, so the answer is unique.
This is widely known as LeetCode 973 (K Closest Points to Origin). It is the standard “top k by a computed key” heap problem.
Constraints for this version: 1 <= k <= len(points) <= 10,000; coordinates between -10,000 and 10,000.
Examples
points |
k |
Result (any order) | Squared distances |
|---|---|---|---|
[[3, 4], [1, -1], [-2, 2], [6, 0]] |
2 |
[[1, -1], [-2, 2]] |
25, 2, 8, 36 |
[[5, 5], [-1, 0]] |
1 |
[[-1, 0]] |
50, 1 |
[[2, 3]] |
1 |
[[2, 3]] |
13 |
Approach 1: brute force
Sort all points by squared distance and take the first k.
def k_closest_sort(points, k):
return sorted(points, key=lambda p: p[0] * p[0] + p[1] * p[1])[:k]
O(n log n) time and O(n) space. Short and correct; a fine first answer, and acceptable when k is close to n.
Approach 2: optimal
Max-heap of size k
Key insight. You only need to remember the best k points so far, and to decide whether a new point belongs among them you compare it with the farthest of those k. A max-heap keyed by distance gives the farthest in O(1). With heapq, use the negated distance as the key.
Walkthrough for the first example, k = 2:
| Point | Distance | Heap after (distances) |
|---|---|---|
| (3, 4) | 25 | 25 |
| (1, -1) | 2 | 25, 2 |
| (-2, 2) | 8 | push gives 25, 8, 2; pop farthest 25: 8, 2 |
| (6, 0) | 36 | push gives 36, 8, 2; pop farthest 36: 8, 2 |
import heapq
def k_closest(points, k):
heap = [] # entries: (-distance, index)
for i, (x, y) in enumerate(points):
d = x * x + y * y
if len(heap) < k:
heapq.heappush(heap, (-d, i))
elif d < -heap[0][0]: # closer than the farthest kept point
heapq.heapreplace(heap, (-d, i))
return [points[i] for _, i in heap]
Storing the index rather than the point list keeps tuple comparisons cheap and avoids comparing lists on ties. heapq.nsmallest(k, points, key=...) implements a similar bounded-heap strategy in the standard library:
def k_closest_nsmallest(points, k):
return heapq.nsmallest(k, points, key=lambda p: p[0] * p[0] + p[1] * p[1])
Quickselect
Key insight. You do not need the k closest points in order, only separated from the rest. Partition the list around a random pivot distance, as quicksort does, then continue only in the side that contains position k.
import random
def k_closest_quickselect(points, k, rng=random.Random(0)):
pts = list(points)
dist = lambda p: p[0] * p[0] + p[1] * p[1]
lo, hi = 0, len(pts) - 1
while lo < hi:
pivot_index = rng.randint(lo, hi)
pts[pivot_index], pts[hi] = pts[hi], pts[pivot_index]
pivot = dist(pts[hi])
store = lo
for i in range(lo, hi): # Lomuto partition
if dist(pts[i]) < pivot:
pts[i], pts[store] = pts[store], pts[i]
store += 1
pts[store], pts[hi] = pts[hi], pts[store]
if store == k - 1 or store == k:
break
if store < k:
lo = store + 1
else:
hi = store - 1
return pts[:k]
Complexity. Heap: O(n log k) time, O(k) space, and it works on streams. Quickselect: O(n) average time with a random pivot, O(n^2) worst case, O(n) space for the copy (O(1) if you may reorder the input).
Tests
def dist(p):
return p[0] * p[0] + p[1] * p[1]
def normalise(result):
return sorted(map(tuple, result))
fns = (k_closest_sort, k_closest, k_closest_nsmallest, k_closest_quickselect)
for fn in fns:
assert normalise(fn([[3, 4], [1, -1], [-2, 2], [6, 0]], 2)) == [(-2, 2), (1, -1)]
assert normalise(fn([[5, 5], [-1, 0]], 1)) == [(-1, 0)]
assert normalise(fn([[2, 3]], 1)) == [(2, 3)]
pts = [[0, 0], [1, 1], [2, 2]]
assert normalise(fn(pts, 3)) == [(0, 0), (1, 1), (2, 2)] # k == n
assert normalise(fn([[1, 1], [1, 1], [5, 5]], 2)) == [(1, 1), (1, 1)] # duplicate points
rng = random.Random(23)
for _ in range(300):
n = rng.randint(1, 40)
pts = [[rng.randint(-50, 50), rng.randint(-50, 50)] for _ in range(n)]
k = rng.randint(1, n)
cutoff = sorted(map(dist, pts))
if k < n and cutoff[k - 1] == cutoff[k]:
continue # skip ties at the boundary (answer not unique)
want = normalise(k_closest_sort(pts, k))
for fn in fns:
assert normalise(fn(pts, k)) == want, (fn.__name__, pts, k)
print("all k closest points tests passed")
Edge cases and pitfalls
- Square roots.
math.sqrtis unnecessary and introduces floating point; compare squared distances, which preserve the order. - Heap direction. A min-heap of size
kkeeps the k farthest points. For the closest, keep a max-heap (negated keys) and evict the farthest. - Comparing lists on ties. Pushing
(distance, point)makes Python compare points when distances tie; that works for lists of ints but fails for objects. Use an index or counter as the second element. - Ties at the boundary. If several points share the k-th distance, the answer is not unique; ask how to break ties.
Where this shows up in data engineering
“Top k by a computed score” is routine: nearest stores to a customer, closest embeddings in a vector search, highest-risk transactions. In SQL it is ORDER BY score LIMIT k, which engines such as Spark execute as a per-partition top-k followed by a merge rather than a full sort. At large scale, nearest-neighbour search uses approximate indexes, but the bounded heap is still how the final top k is collected.
Progress is saved in this browser only. No account needed.