Skip to content
BytePatterns

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 saves n square roots and avoids floating-point noise.
  • Python's heapq is 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 size k. It works on a stream, because it never needs to see more than one point at a time.
  • Sort: O(n log n) time and O(n) space. Simplest to write; when k is close to n, 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 n points. Heapifying everything and popping k times is O(n + k log n) time but O(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 k survivors: 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.