A k-d tree is a binary search tree for points that have more than one coordinate. Each internal node holds a point and a splitting axis; everything in its left subtree lies on one side of an axis-aligned plane through that point, and everything in its right subtree lies on the other. That single idea, cutting space in half one coordinate at a time, gives logarithmic-ish nearest-neighbour search in two or three dimensions, fast rectangular range queries, and a structure that is still the default for point clouds, geospatial lookups, collision detection and low-dimensional machine-learning features.

It also has a well-known failure: as dimensions grow, the pruning that makes it fast stops working and the tree degrades into an expensive linear scan. This article builds the structure from first principles, traces a search by hand, gives working Python for construction, nearest neighbour, k-NN and range search, and covers what separates a textbook tree from a production one.

Advertisement

The idea: a search tree whose key changes with depth

In an ordinary binary search tree every node compares the same key. A k-d tree, introduced by Jon Bentley in 1975, compares a different coordinate at each level. With two-dimensional points, the root splits on x, its children split on y, their children on x again, and so on. The comparison at each node is a question about one coordinate only: is the point's x less than 7? That makes each node an axis-aligned plane (a line in 2-D), and each subtree a rectangular cell of space bounded by the planes of its ancestors.

Two properties follow. First, a subtree's cell is known without storing it: walk down from the root and each comparison tightens one side of a box. Second, for any query point, the distance from the query to a node's splitting plane is a lower bound on the distance to every point in the far subtree, because any path from the query into that half-space has to cross the plane. Every speed-up the structure offers comes from that lower bound.

The six-point example: the plane is cut by alternating axes, and the tree records the cutsx = 7y = 4y = 6(2,3)(5,4)(9,6)(4,7)(8,1)(7,2)query (9,2)(7,2)split x(5,4)split y(9,6)split y(2,3)(4,7)(8,1)nearestx < 7x >= 7Search for (9,2) visits (7,2), (9,6), (8,1).The dashed circle (radius sqrt 2) never crosses y = 6 or x = 7,so the other three points are pruned without being examined.Coordinates are the classic textbook example; the tree is exactly what the build code on this page produces
Left: the plane cut first at x = 7, then each half cut on y. Right: the same cuts as a tree. A nearest-neighbour search for (9,2) touches only the three highlighted nodes.

Construction by median split

To keep the tree balanced, split each cell at the median of the points along the chosen axis. The median point becomes the node, points before it in sorted order go left, points after it go right, and the procedure recurses with the next axis. The result has depth about log2 n, which bounds the cost of descending to a leaf.

from dataclasses import dataclass
import heapq

@dataclass
class Node:
    point: tuple
    axis: int
    left: "Node | None" = None
    right: "Node | None" = None

def build(points, depth=0):
    if not points:
        return None
    k = len(points[0])
    axis = depth % k                                  # round-robin axis choice
    points = sorted(points, key=lambda p: p[axis])
    mid = len(points) // 2
    return Node(points[mid], axis,
                build(points[:mid], depth + 1),
                build(points[mid + 1:], depth + 1))

def sqdist(a, b):
    return sum((x - y) ** 2 for x, y in zip(a, b))

Sorting at every level costs O(n log n) per level and O(n log^2 n) overall. Two standard improvements bring it down to O(n log n): pre-sort the points once on every axis and split the sorted lists as you recurse, or replace the sort with a linear-time selection such as quickselect (the partition step from quicksort) that only places the median and partitions around it. In C++ that is std::nth_element; in NumPy it is np.argpartition.

Duplicates need care: points equal to the median coordinate can land on either side, so the search must treat the plane as belonging to both halves, which the code below does.

Advertisement

Nearest-neighbour search and the pruning rule

Search descends the way an insertion would, recording the closest point seen so far. On the way back up, at each node it asks one question: could the far subtree contain anything closer than the current best? The far subtree lies entirely beyond the splitting plane, so if the squared distance from the query to the plane is already at least the best squared distance, the whole subtree is skipped.

def nearest(node, q, best=None):
    if node is None:
        return best
    d = sqdist(node.point, q)
    if best is None or d < best[0]:
        best = (d, node.point)
    diff = q[node.axis] - node.point[node.axis]
    near, far = (node.left, node.right) if diff < 0 else (node.right, node.left)
    best = nearest(near, q, best)          # go toward the query first
    if diff * diff < best[0]:              # does the best-distance ball cross the plane?
        best = nearest(far, q, best)
    return best

Visiting the near side first finds a good candidate early, which shrinks the ball so pruning succeeds more often on the way back. Squared distances avoid a square root per node.

This is branch-and-bound, and its guarantees are probabilistic. Friedman, Bentley and Finkel showed in 1977 that in a fixed low dimension, with well-behaved data, the expected nodes visited grow logarithmically with n. The worst case is still linear.

Worked example: six points, one query

Take the points (2,3), (5,4), (9,6), (4,7), (8,1) and (7,2). Sorted by x, the median at index 3 is (7,2), so it becomes the root. The left cell holds (2,3), (5,4), (4,7); sorted by y its median is (5,4), with (2,3) below and (4,7) above. The right cell holds (8,1) and (9,6); sorted by y the median at index 1 is (9,6), with (8,1) as its left child. That is the tree in the diagram.

Now search for the point nearest to (9,2).

  1. At the root (7,2) the squared distance is 4, so the best is (7,2) at 4. The query's x is 9, which is not less than 7, so the near side is the right subtree.
  2. At (9,6) the squared distance is 0 + 16 = 16, no improvement. This node splits on y, and the query's y of 2 is below 6, so the near side is its left child.
  3. At (8,1) the squared distance is 1 + 1 = 2. The best becomes (8,1) at 2. It is a leaf, so return.
  4. Back at (9,6): the distance to the plane y = 6 is 4, squared 16, which is not less than 2. The empty right side is pruned.
  5. Back at the root: the distance to the plane x = 7 is 2, squared 4, which is not less than 2. The entire left subtree, three points, is pruned without being examined.

The answer is (8,1) at distance sqrt 2, found with three distance computations instead of six. At a million 3-D points the same rule skips almost every node.

k nearest neighbours and range queries

Finding the k nearest points changes one thing: the pruning radius becomes the distance to the k-th best candidate, kept in a bounded max-heap. Until the heap is full, nothing can be pruned. Python's heapq is a min-heap, so the code stores negated distances; the heap operations article explains why push and replace are logarithmic.

def knn(root, q, k):
    heap = []                                    # (-squared_distance, point), max-heap
    def visit(node):
        if node is None:
            return
        d = sqdist(node.point, q)
        if len(heap) < k:
            heapq.heappush(heap, (-d, node.point))
        elif d < -heap[0][0]:
            heapq.heapreplace(heap, (-d, node.point))
        diff = q[node.axis] - node.point[node.axis]
        near, far = (node.left, node.right) if diff < 0 else (node.right, node.left)
        visit(near)
        if len(heap) < k or diff * diff < -heap[0][0]:
            visit(far)
    visit(root)
    return sorted((-nd, p) for nd, p in heap)

def range_query(node, lo, hi, out):
    # every point p with lo[i] <= p[i] <= hi[i] on all axes
    if node is None:
        return out
    p, a = node.point, node.axis
    if all(l <= x <= h for x, l, h in zip(p, lo, hi)):
        out.append(p)
    if lo[a] <= p[a]:
        range_query(node.left, lo, hi, out)
    if p[a] <= hi[a]:
        range_query(node.right, lo, hi, out)
    return out

An orthogonal range query descends into a child only when the query box overlaps that child's side of the plane. For a balanced tree in k dimensions the worst case is O(n^(1 - 1/k) + m), where m is the number of points reported; in two dimensions that is O(sqrt n + m). Radius queries, such as every point within 5 metres, work the same way with a ball instead of a box, and are what scipy.spatial.KDTree.query_ball_point answers.

Why it collapses in high dimensions

The pruning test compares the distance to one plane against the distance to the best point. In two dimensions those are similar in scale, so the test often succeeds. In fifty dimensions they are not: the distance to any single coordinate plane is one term of a fifty-term sum, while the distance to the nearest neighbour is the whole sum. The ball around the query crosses almost every plane, and almost no subtree gets pruned.

A second effect compounds it. In high dimensions, distances between random points concentrate: the nearest and farthest neighbours of a query end up at nearly the same distance, so there is little contrast for any exact method to exploit. A common rule of thumb is that a k-d tree helps only when n is much larger than 2^k; with 2^k cells needed just to split once on every axis, a 30-dimensional tree never gets deep enough to split on most coordinates.

This is why embeddings of hundreds of dimensions use approximate indexes: graphs such as HNSW, partition-and-quantise schemes such as IVF-PQ and disk-resident graphs such as DiskANN trade exactness for cost that does not explode with dimension. For small collections, brute force on a GPU is often fastest, because a matrix multiply is perfectly regular work.

From textbook to production

The tree above follows a pointer per point, which is slow in practice. Production implementations differ in four ways.

  • Bucketed leaves. Stop splitting when a cell holds a small number of points, often somewhere between 8 and 64, and scan the bucket linearly. Scanning a contiguous bucket is cheap and vectorises well; descending through near-empty nodes is not. SciPy exposes this as the leafsize argument.
  • Flat array layout. Reorder the points so each subtree occupies a contiguous slice of one array, and store nodes as an array of split values and child indices. There are no per-node allocations, the tree serialises with a memory copy, and leaves are cache-friendly.
  • Smarter split rules. Round-robin axes waste splits on coordinates with little spread. Splitting on the axis of widest spread adapts to the data. The sliding-midpoint rule, which SciPy documents for its tree, splits a cell at its geometric midpoint and slides the plane to the nearest point if one side would be empty; it avoids long thin cells that hurt search.
  • Approximate search. Pruning with a slightly shrunken ball, prune when diff squared times (1 + eps) squared is at least the best, returns a point within a factor (1 + eps) of the true nearest and visits far fewer nodes. SciPy's query exposes this as eps.

Updates, metrics and data hazards

Inserting a point is easy: descend and attach a leaf. Doing it repeatedly unbalances the tree, and deletion is harder, because replacing an internal node means searching subtrees that split on other axes. Most systems rebuild in the background when enough changes accumulate, answering queries against the tree plus a small linear buffer of new points meanwhile. The logarithmic method keeps static trees of sizes 1, 2, 4 and so on and merges them like a binary counter, giving amortised O(log^2 n) insertion.

The pruning bound holds for Euclidean distance and other Minkowski norms such as Manhattan and Chebyshev, but not for cosine similarity directly; normalise vectors to unit length first, after which Euclidean ranking matches cosine ranking. Standardise features too, or one in millimetres will dominate one in metres.

One data hazard causes silent wrong answers: a NaN coordinate makes every comparison false, so the point lands on an arbitrary side and the search prunes wrongly. Reject NaNs at ingest.

Failure modes and trade-offs

  • Silent linear scans. In 20 or more dimensions a tree may visit most leaves while looking fast in small tests. Measure nodes visited per query, not just latency.
  • Clustered data. Many points with identical coordinates on the split axis produce empty or degenerate cells. Widest-spread or sliding-midpoint splits help.
  • Stale trees. A tree built from yesterday's points answers with yesterday's neighbours.
OptionBest forCost
k-d treeExact queries in roughly 2 to 15 dimensions; range and radius queriesPruning fades with dimension; updates need rebuilds
Ball treeModerate dimensions and non-axis-aligned clustersMore expensive to build; still exact and still dimension-sensitive
HNSW or IVF-PQHundreds of dimensions, millions of vectorsApproximate; recall must be measured
Brute forceSmall collections, or batches on a GPULinear per query, but simple and exact

What to do next

  1. Run the code on this page against the six-point example and confirm the search for (9,2) computes exactly three distances.
  2. Build a tree over one million random 3-D points and log nodes visited per query; repeat at 10, 20 and 50 dimensions and watch pruning disappear.
  3. Switch to bucketed leaves and a widest-spread split, and compare build time and query latency against the textbook version.
  4. Standardise features, normalise vectors if you use cosine, and reject NaNs before any tree is built.
  5. Before choosing an index for embeddings, benchmark brute force and an approximate index at your real dimension and recall target.
  6. If your data changes, decide your rebuild trigger and keep a small linear buffer for points added since the last build.
Key takeaway: A k-d tree splits space one coordinate at a time, and its speed comes from a single lower bound: nothing beyond a splitting plane can be closer than the plane itself. Build it with median or sliding-midpoint splits and bucketed leaves, search near side first, and prune the far side whenever the plane is farther than your current best. It is the right exact index for low-dimensional points such as locations, point clouds and small feature vectors. In high dimensions the pruning stops firing, so measure nodes visited, and move to approximate indexes or brute force when the tree turns into a scan.