K Closest Points to Origin: Max-Heap vs Sort vs Quickselect
8 min readBytePatterns
Why a heap of size k keeps the farthest point on top, when sorting or quickselect wins instead, and why squared distance is enough. Every version tested.
"Return the k points closest to the origin" is one of the cleanest heap questions there is, and one of the easiest to get subtly wrong. Most people reach for a min-heap because they want the smallest distances. The version that scales uses a max-heap, and the reason it does is the whole idea behind every "top k" problem.
The problem it solves
You have n points on a plane and a number k. Return the k points with the smallest Euclidean distance to (0, 0), in any order. The same shape shows up everywhere: the nearest drivers to a rider, the closest stores to a postcode, the k embeddings most similar to a query when you are not using an index.
The obvious answer is to sort all n points by distance and take the first k. That is correct and often fine. It stops being fine when n is much larger than k, or when the points arrive as a stream and you cannot hold them all.
The intuition
Keep a shortlist of exactly k candidates. When a new point arrives, the only question is: is it closer than the worst point on the shortlist? If it is, the worst one leaves and the new one takes its place. If it is not, the new point is ignored.
So the shortlist needs one operation to be instant: "show me the farthest point I am holding". That is a max-heap keyed by distance. A min-heap would give you the closest point instantly, which is the one you never want to throw away.
Two details make it cheap:
- Squared distance is enough. The square root is increasing, so
x² + y²orders points exactly as the true distance does. Skipping it savesnsquare roots and avoids floating-point noise. - Python's
heapqis a min-heap. Push the negated distance and the largest distance surfaces as the smallest key.
Watch it run
The animation feeds the lesson's four points into a heap capped at two. Each point enters, the farthest survivor floats to the root, and as soon as the heap holds three, the root is evicted. The point (5, 8), at squared distance 89, is pushed and thrown straight back out. What remains are (-2, 2) and (0, 1).
K Closest Points
Step 1 of 8
Four points, and only the 2 nearest matter. Distances are squared — the square root would not change any comparison.
The same interactive animation as the lesson — step through it with the controls.
The code
The heap version compares against the root before pushing, so a point that cannot make the shortlist costs one comparison and no heap work. heapreplace pops the root and pushes the new item in a single sift:
import heapq
def k_closest_heap(points, k):
heap = [] # max-heap by distance, via negation
for x, y in points:
d = x * x + y * y # squared: same order, no sqrt
if len(heap) < k:
heapq.heappush(heap, (-d, x, y))
elif d < -heap[0][0]: # closer than the worst one kept
heapq.heapreplace(heap, (-d, x, y)) # pop the root and push in one step
return [(x, y) for _, x, y in heap]
pts = [(1, 3), (-2, 2), (5, 8), (0, 1)]
print(sorted(k_closest_heap(pts, 2))) # [(-2, 2), (0, 1)]
Two shorter alternatives. Sorting is the baseline; heapq.nsmallest does the bounded-heap trick internally for you:
def k_closest_sort(points, k):
return sorted(points, key=lambda p: p[0] * p[0] + p[1] * p[1])[:k]
def k_closest_heapq(points, k):
return heapq.nsmallest(k, points, key=lambda p: p[0] * p[0] + p[1] * p[1])
print(k_closest_sort(pts, 2), k_closest_heapq(pts, 2))
# [(0, 1), (-2, 2)] [(0, 1), (-2, 2)]
The third option is quickselect: partition around a random pivot distance, keep only the side that contains position k, and repeat. When it stops, the first k slots hold the answer, unsorted:
import random
def k_closest_select(points, k):
pts = list(points)
dist = lambda p: p[0] * p[0] + p[1] * p[1]
lo, hi = 0, len(pts) - 1
while lo < hi:
pivot = dist(pts[random.randint(lo, hi)])
i, j = lo, hi
while i <= j: # Hoare-style partition around pivot
while dist(pts[i]) < pivot:
i += 1
while dist(pts[j]) > pivot:
j -= 1
if i <= j:
pts[i], pts[j] = pts[j], pts[i]
i, j = i + 1, j - 1
if k - 1 <= j:
hi = j # the k-th smallest is on the left
elif k - 1 >= i:
lo = i # ... or on the right
else:
break # it sits in the pivot's band
return pts[:k]
Ties make the answer non-unique — two points at the same distance may both qualify for the last slot — so the check compares distances, not points. On 2,000 random inputs with small coordinates, all four versions return exactly the k smallest distances, and only points that were in the input:
random.seed(11)
ok = True
ties = 0
for _ in range(2000):
n = random.randint(1, 30)
pts = [(random.randint(-6, 6), random.randint(-6, 6)) for _ in range(n)]
k = random.randint(1, n)
d = lambda p: p[0] * p[0] + p[1] * p[1]
want = sorted(d(p) for p in pts)[:k] # brute force: the k smallest distances
for f in (k_closest_heap, k_closest_sort, k_closest_heapq, k_closest_select):
got = f(pts, k)
ok &= sorted(d(p) for p in got) == want and len(got) == k
ok &= all(got.count(p) <= pts.count(p) for p in got) # only real input points
ties += len(set(map(d, pts))) < n
print(ok, ties) # True 1633
1,633 of the 2,000 inputs contained a tie, so the tie handling was exercised, not assumed.
The complexity
- Bounded heap:
O(n log k)time,O(k)extra space. Every point costs at most one push or replace on a heap of sizek. It works on a stream, because it never needs to see more than one point at a time. - Sort:
O(n log n)time andO(n)space. Simplest to write; whenkis close ton, it is as good as anything else. - Quickselect:
O(n)on average,O(n²)in the worst case, and it needs all points in memory because it rearranges them. A random pivot makes the worst case unlikely rather than impossible.
A useful rule: stream or k much smaller than n → heap; everything already in an array and speed matters → quickselect; otherwise, sort and move on.
Where it goes wrong
- Using a min-heap of all
npoints. Heapifying everything and poppingktimes isO(n + k log n)time butO(n)memory — it loses the streaming property that makes the heap answer worth giving. - Forgetting to negate. With
heapq, pushing the positive distance keeps the closest point at the root, and the eviction throws away exactly the points you wanted. - Taking square roots. Not wrong, just slower, and floating-point results can make two equal distances compare unequal.
- Returning the heap as "sorted". A heap is only ordered at the root. If the caller needs the points in order, sort the
ksurvivors:O(k log k)extra.
How to say it in an interview
"I keep a max-heap of size k keyed by squared distance — squared because the square root doesn't change the order. For each point, if the heap has fewer than k items I push it; otherwise, if it's closer than the root, I replace the root. The root is always the worst point I'm keeping, so it's the one to evict. That's O(n log k) time and O(k) space, and it works on a stream. If all the points are in memory, quickselect gets O(n) on average."
The same keep-the-worst-on-top trick powers top k elements, and the partition step behind quickselect is the one from quick sort.