PyInfo
Computational Geometry contents

Closest Pair of Points

Find the two nearest points among n in O(n log n) by divide and conquer, and in expected O(n) with randomized grid algorithms.

Advanced10 min readgeometryclosest pairdivide and conquerrandomized algorithmsgrid hashing

Read first: Basic Geometry: Vectors, Dot and Cross Products, Heaps, Deques and Bisect

Given nn points in the plane, find the two whose Euclidean distance is smallest. Checking all pairs takes O(n2)O(n^2); the algorithm of Shamos and Hoey (1975) does it in O(nlog⁡n)O(n\log n), which is optimal in the decision-tree model. There are also simple randomized algorithms with expected linear time.

Throughout, we work with squared distances in integers to avoid square roots and rounding errors.

python
def dist2(a, b):
    return (a[0] - b[0]) ** 2 + (a[1] - b[1]) ** 2

def closest_pair_bruteforce(points):
    """O(n^2) reference: (squared distance, i, j)."""
    best = None
    for i in range(len(points)):
        for j in range(i + 1, len(points)):
            d = dist2(points[i], points[j])
            if best is None or d < best[0]:
                best = (d, i, j)
    return best

Divide and conquer

Sort the points by xx (ties by yy). Split the sorted list in the middle into A1A_1 and A2A_2, solve each half recursively and let h=min⁡(h1,h2)h = \min(h_1, h_2) be the better of the two answers. The only pairs we might still miss are those with one point in each half; both such points must be within distance hh of the dividing vertical line.

For this strip B={p:∣xp−xm∣<h}B = \{p : |x_p - x_m| < h\}, sort the points by yy and compare every point only with the previous points whose yy differs by less than hh. Surprisingly, there are only O(1)O(1) such points for each point:

Why at most 7 comparisons. The candidate points for pip_i lie in a 2h×h2h\times h rectangle. Split it into the two h×hh\times h squares belonging to the two halves. Inside each half any two points are at least hh apart, so each h×hh\times h square holds at most 4 points (dividing it into four h/2h/2-squares, each with diagonal h/2<hh/\sqrt2 < h, at most one point per sub-square). So at most 88 points, one of which is pip_i itself.

To keep the merge step linear, the recursion returns its points sorted by yy, like merge sort; the strip is then extracted in yy order without sorting again.

python
def closest_pair(points):
    """Divide and conquer, O(n log n). Returns (squared distance, point a, point b)."""
    pts = sorted(points)                              # by x, then y
    best = [float("inf"), None, None]

    def update(a, b):
        d = dist2(a, b)
        if d < best[0]:
            best[0], best[1], best[2] = d, a, b

    def rec(lo, hi):                                  # solves pts[lo:hi] and leaves it sorted by y
        if hi - lo <= 3:
            for i in range(lo, hi):
                for j in range(i + 1, hi):
                    update(pts[i], pts[j])
            pts[lo:hi] = sorted(pts[lo:hi], key=lambda p: p[1])
            return
        mid = (lo + hi) // 2
        mid_x = pts[mid][0]
        rec(lo, mid)
        rec(mid, hi)
        pts[lo:hi] = sorted(pts[lo:hi], key=lambda p: p[1])      # merging two sorted runs: linear for Timsort
        strip = []
        for i in range(lo, hi):
            p = pts[i]
            if (p[0] - mid_x) ** 2 < best[0]:
                for q in reversed(strip):
                    if (p[1] - q[1]) ** 2 >= best[0]:
                        break
                    update(p, q)
                strip.append(p)

    rec(0, len(pts))
    return tuple(best)

assert closest_pair([(0, 0), (10, 10), (3, 4), (11, 12), (0, 5)])[0] == 5          # the closest pair is (10, 10)-(11, 12)
assert closest_pair([(0, 0), (3, 4)])[0] == 25
assert closest_pair([(1, 1), (5, 5), (1, 1)])[0] == 0                                # duplicates

Python's sorted (Timsort) detects that the slice consists of two already sorted runs and merges them in linear time, so the recursion has the intended T(n)=2T(n/2)+O(n)T(n) = 2T(n/2) + O(n) cost.

Randomized algorithms with a grid

Rabin / Lipton: sample, then use a grid

Cut the plane into squares of side dd. If dd is at least the true minimum distance, then any pair of points at distance at most dd lies in the same or in adjacent squares, so it is enough to compare points in the same and neighbouring squares. The cost is Θ(∑ni2)\Theta(\sum n_i^2) where nin_i are the numbers of points in the non-empty squares.

To choose dd: sample nn random pairs and let dd be the smallest distance found. It can be shown that then E[∑ni2]≤16n\mathbb E\big[\sum n_i^2\big] \le 16n, so the whole algorithm is expected linear. Duplicated points must be handled first (with a hash set), because a grid with d=0d = 0 is meaningless.

python
import math
import random
from collections import defaultdict

def closest_pair_grid(points, seed=1):
    """Expected O(n): sample a distance, bucket the points in a grid, compare neighbouring cells."""
    n = len(points)
    assert n >= 2
    seen = {}
    for i, p in enumerate(points):
        if p in seen:
            return 0, points[seen[p]], p                 # a duplicate is a pair at distance 0
        seen[p] = i

    rnd = random.Random(seed)
    best = [dist2(points[0], points[1]), points[0], points[1]]

    def consider(a, b):
        d = dist2(a, b)
        if d < best[0]:
            best[0], best[1], best[2] = d, a, b

    for _ in range(n):
        i, j = rnd.sample(range(n), 2)
        consider(points[i], points[j])

    d = math.isqrt(best[0]) + 1                          # a cell side that is >= the sampled distance
    grid = defaultdict(list)
    for p in points:
        grid[(p[0] // d, p[1] // d)].append(p)
    for (cx, cy), cell in grid.items():
        for i in range(len(cell)):
            for j in range(i + 1, len(cell)):
                consider(cell[i], cell[j])
        for dx, dy in ((1, -1), (1, 0), (1, 1), (0, 1)):     # each pair of neighbouring cells once
            other = grid.get((cx + dx, cy + dy))
            if other:
                for p in cell:
                    for q in other:
                        consider(p, q)
    return tuple(best)

Incremental with a shrinking grid

A different algorithm is easier to analyze. Shuffle the points; let δ\delta be the distance of the first two. Keep the points seen so far in a grid of side about δ/2\delta/2. Insert points one at a time: look at the 5×55\times5 block of cells around the new point (any point closer than δ\delta must be within two cells). If a closer point is found, δ\delta decreases and we rebuild the grid from the first ii points; otherwise just insert the point.

The rebuild happens at step ii only if pip_i belongs to the closest pair of the first ii points, which for a random order has probability at most 2/i2/i; a rebuild costs O(i)O(i), so the total expected cost is ∑ii⋅2i=O(n)\sum_i i\cdot\frac 2i = O(n).

python
def closest_pair_incremental(points, seed=1):
    """Expected O(n) by inserting points in random order into a grid that is rebuilt when the minimum improves."""
    pts = list(points)
    random.Random(seed).shuffle(pts)
    best = [dist2(pts[0], pts[1]), pts[0], pts[1]]

    def cell_size():
        r = math.isqrt(best[0])
        if r * r < best[0]:
            r += 1                                       # r = ceil(sqrt(best))
        return max(1, (r + 1) // 2)                      # >= delta / 2, so neighbours are within 2 cells

    def build(count):
        size = cell_size()
        grid = defaultdict(list)
        for p in pts[:count]:
            grid[(p[0] // size, p[1] // size)].append(p)
        return size, grid

    size, grid = build(2)
    for i in range(2, len(pts)):
        p = pts[i]
        if best[0] == 0:
            break
        cx, cy = p[0] // size, p[1] // size
        found = None
        for gx in range(cx - 2, cx + 3):
            for gy in range(cy - 2, cy + 3):
                for q in grid.get((gx, gy), ()):
                    d = dist2(p, q)
                    if d < best[0] and (found is None or d < found[0]):
                        found = (d, p, q)
        if found:
            best[:] = found
            size, grid = build(i + 1)                    # the minimum shrank: rebuild with the points so far
        else:
            grid[(cx, cy)].append(p)
    return tuple(best)

Testing

All three algorithms must report the same squared distance as the brute force, and the returned points must be at that distance:

python
rnd = random.Random(5)
for _ in range(1500):
    n = rnd.randint(2, 40)
    r = rnd.choice([5, 20, 1000])
    pts = [(rnd.randint(-r, r), rnd.randint(-r, r)) for _ in range(n)]
    expected = closest_pair_bruteforce(pts)[0]
    for solver in (closest_pair, closest_pair_grid, closest_pair_incremental):
        d, a, b = solver(pts)
        assert d == expected == dist2(a, b), (solver.__name__, pts)
        assert a in pts and b in pts

On a larger input only the fast algorithms are practical (the quadratic reference would need 200 million distance computations); the three independent algorithms must agree:

python
import time

big = [(rnd.randint(0, 10 ** 7), rnd.randint(0, 10 ** 7)) for _ in range(20_000)]
start = time.perf_counter()
d, a, b = closest_pair(big)
elapsed = time.perf_counter() - start
assert d == closest_pair_grid(big)[0] == closest_pair_incremental(big)[0]
assert elapsed < 10

Generalization: the triangle with minimum perimeter

The same divide-and-conquer applies to finding three points with the smallest sum of pairwise distances. After solving both halves, let minperminper be the best perimeter found; in a triangle of perimeter at most minperminper the longest side is at most minper/2minper/2, so take the strip of width minper/2minper/2 and check the triangles inside it that can improve the answer.

Practice problems

This article is a Python adaptation of “Finding the nearest pair of points” from cp-algorithms.com, licensed under CC BY-SA 4.0. The text was condensed and rewritten and the C++ code was reimplemented in Python; this adaptation is shared under the same license.