A k-d tree splits space with axis-aligned planes so that a nearest-neighbour or range query can skip whole regions without looking at their points. The textbook version, with one point per node and a cycling split axis, is covered step by step in the k-d tree deep dive. This page picks up where that one stops and is about building one that is fast and correct in practice: pruning with bounding boxes instead of splitting planes, searching best-first, counting points in a rectangle without visiting them, and measuring the point at which the tree stops paying for itself.

You will get a compact implementation that was tested against a brute-force oracle on hundreds of random inputs, including grids full of duplicate points, a worked query traced by hand, measured numbers showing how pruning decays with dimension, and a checklist for deciding whether a k-d tree is the right index for your data at all.

The model in two paragraphs

Every node of a k-d tree owns a set of points. An internal node chooses an axis and a split value and hands the points below the split to its left child and the rest to its right child; a leaf stores a small bucket of points and is scanned linearly. Building by the median on each node keeps the tree balanced, so its depth is about log2(n / leaf_size).

A query is correct as long as it never discards a node that could contain an answer. Search order, bucket size and the discard bound affect only speed, provided the bound never overestimates the distance to anything inside the node.

Pruning: splitting planes versus bounding boxes

query qleft cell boxright cell boxsplit plane x = sPlane test: distance to x = sis small, so both sides openBox test: distance to thetight bounding box of thepoints actually in the cellEmpty margins inside a cellmake the box test strictlytighter, so more cells pruned
Two lower bounds for the same cell. The splitting plane is close to q, but the points on each side occupy smaller boxes, which are farther away.

The classic pruning test compares the current best distance with the distance from the query to the splitting plane. It is a valid lower bound, but a loose one: it knows only one coordinate of the cell's boundary. If the cell's points are clustered far from the plane, the test still opens the cell.

A tighter bound comes from storing, for each node, the bounding box of the points it actually contains. The squared distance from a query q to a box is the sum, over dimensions, of the squared gap between q and the box on that axis, with zero for any axis where q lies inside the box's range. Because every point in the node lies inside the box, no point can be closer than that distance, so it is a valid lower bound. And because the box is never larger than the cell, it is at least as tight as the plane test, often far tighter near the edges of data where cells contain large empty margins.

Boxes cost 2k numbers per node and a pass over the points during build. They also make ties harmless: points equal to the split value may land on either side, and the boxes simply describe wherever they landed.

A tested implementation

The implementation below stores the tree in parallel arrays rather than node objects, splits on the axis of widest spread at the median, keeps bucket leaves, and stores a tight box for each node. Nearest-neighbour search is best-first: a priority queue holds nodes ordered by their box distance, and the search stops as soon as the closest unopened box is no closer than the k-th best point found. Range counting uses the boxes in a second way, described below.

import heapq


class KDTree:
    """Static k-d tree in flat arrays. Each node owns pts[lo:hi] and a bounding box."""

    def __init__(self, points, leaf_size=16):
        self.pts = [tuple(p) for p in points]
        self.k = len(self.pts[0])
        self.leaf_size = leaf_size
        # parallel node arrays: slice, split axis, split value, children, box
        self.lo, self.hi, self.axis, self.split = [], [], [], []
        self.left, self.right, self.bmin, self.bmax = [], [], [], []
        self._build()

    def _new_node(self, lo, hi):
        sl = self.pts[lo:hi]
        self.lo.append(lo); self.hi.append(hi)
        self.bmin.append(tuple(min(p[d] for p in sl) for d in range(self.k)))
        self.bmax.append(tuple(max(p[d] for p in sl) for d in range(self.k)))
        self.axis.append(-1); self.split.append(0.0)
        self.left.append(-1); self.right.append(-1)
        return len(self.lo) - 1

    def _build(self):
        stack = [self._new_node(0, len(self.pts))]
        while stack:
            n = stack.pop()
            lo, hi = self.lo[n], self.hi[n]
            if hi - lo <= self.leaf_size:
                continue                      # leaf: scanned linearly at query time
            spread = [self.bmax[n][d] - self.bmin[n][d] for d in range(self.k)]
            ax = max(range(self.k), key=spread.__getitem__)
            if spread[ax] == 0:
                continue                      # all points identical: keep as a big leaf
            mid = (lo + hi) // 2
            seg = sorted(self.pts[lo:hi], key=lambda p: p[ax])   # quickselect in production
            self.pts[lo:hi] = seg
            self.axis[n], self.split[n] = ax, seg[mid - lo][ax]
            self.left[n] = self._new_node(lo, mid)
            self.right[n] = self._new_node(mid, hi)
            stack += [self.left[n], self.right[n]]

    def _box_dist2(self, n, q):
        """Squared distance from q to node n's bounding box (0 if q is inside)."""
        s = 0.0
        for d in range(self.k):
            v = q[d]
            if v < self.bmin[n][d]:
                s += (self.bmin[n][d] - v) ** 2
            elif v > self.bmax[n][d]:
                s += (v - self.bmax[n][d]) ** 2
        return s

    def knn(self, q, k=1):
        """k nearest points to q as (dist2, point), nearest first."""
        best = []                             # max-heap of (-dist2, point)
        visited = 0
        frontier = [(self._box_dist2(0, q), 0)]   # min-heap of (box dist2, node)
        while frontier:
            bd, n = heapq.heappop(frontier)
            if len(best) == k and bd >= -best[0][0]:
                break                         # every remaining box is farther: done
            visited += 1
            if self.left[n] == -1:
                for p in self.pts[self.lo[n]:self.hi[n]]:
                    d2 = sum((a - b) ** 2 for a, b in zip(p, q))
                    if len(best) < k:
                        heapq.heappush(best, (-d2, p))
                    elif d2 < -best[0][0]:
                        heapq.heapreplace(best, (-d2, p))
                continue
            for c in (self.left[n], self.right[n]):
                heapq.heappush(frontier, (self._box_dist2(c, q), c))
        self.last_visited = visited
        return sorted((-nd, p) for nd, p in best)

    def range_count(self, qmin, qmax):
        """Number of points p with qmin[d] <= p[d] <= qmax[d] for every d."""
        total, stack = 0, [0]
        while stack:
            n = stack.pop()
            bmin, bmax = self.bmin[n], self.bmax[n]
            if any(bmax[d] < qmin[d] or bmin[d] > qmax[d] for d in range(self.k)):
                continue                      # disjoint: prune
            if all(qmin[d] <= bmin[d] and bmax[d] <= qmax[d] for d in range(self.k)):
                total += self.hi[n] - self.lo[n]   # fully inside: count without visiting
                continue
            if self.left[n] == -1:
                total += sum(all(qmin[d] <= p[d] <= qmax[d] for d in range(self.k))
                             for p in self.pts[self.lo[n]:self.hi[n]])
                continue
            stack += [self.left[n], self.right[n]]
        return total

Two details matter for correctness. The build stops splitting when every point in a node is identical, because no axis has spread; without that guard a grid of duplicates would recurse forever. And the search compares squared distances throughout, which avoids square roots and keeps the bound and the candidates in the same units. Sorting each node's slice makes the build O(n log² n); replacing it with a linear-time selection such as C++ std::nth_element or NumPy argpartition gives O(n log n), which is what libraries do.

Worked example: tracing one query

Take six points in the plane: A(2,3), B(5,4), C(9,6), D(4,7), E(8,1) and F(7,2), with a leaf size of 2. The root's box spans x from 2 to 9 and y from 1 to 7, so x has the wider spread. Sorted by x, the left half is A, D, B and the right half is F, E, C. The left node's box spans x 2 to 5 and y 3 to 7, so it splits on y into a leaf holding A and a leaf holding B and D. The right node's box spans x 7 to 9 and y 1 to 6, so it splits on y into a leaf holding E and a leaf holding F and C.

StepNode openedBox distance²ActionBest so far
1root0push both children (each at distance² 1)none
2left internal1push leaf A (20) and leaf B, D (1)none
3right internal1push leaf E (20) and leaf F, C (1)none
4leaf F, C1F and C both at distance² 10F, 10
5leaf B, D1B at 2, D at 8B, 2
6leaf E (A is also at 20)2020 is not below 2: stopB, 2

The query is q = (6, 5) with k = 1. The answer is B at distance √2. Two of the four leaves were never scanned, and the search ended at the frontier rather than by exhausting it. Notice step 4: the first leaf opened was not the one containing the answer. Best-first search does not guarantee that, but it opens the most promising region next, and the first good candidate it finds immediately shrinks the bound for everything else.

Depth-first or best-first

The textbook search is depth-first: descend to the leaf containing q, then unwind and check siblings, using only a stack of depth O(log n). Best-first search, which the code uses, keeps a priority queue of unopened nodes ordered by lower bound and can stop the moment that bound exceeds the k-th best distance. Both are exact. Best-first opens fewer nodes but its queue can grow large when pruning is weak.

Best-first also offers a practical approximation knob: stop after opening, say, 32 leaves and return the best found. You lose exactness but gain bounded latency, the same trade graph-based indexes make.

Range counting and the containment shortcut

For an orthogonal range query, a box with a lower and upper bound on each axis, a node relates to the query box in exactly one of three ways. If the node's box is disjoint from the query, prune it. If the node's box lies entirely inside the query, every point below it matches, so for a count you add the node's size without visiting anything. Only partially overlapping nodes are opened.

That containment shortcut is what makes counting cheap. For a balanced k-d tree the number of partially overlapping nodes is O(n^(1−1/k)), a classic result; reporting the matching points adds the output size m on top, but counting does not. In two dimensions that is about √n nodes, so counting points inside a map viewport over a million points touches on the order of a thousand nodes rather than a million points. Store per-node aggregates other than the count, such as a sum or maximum of a value attached to each point, and the same traversal answers questions like total sales inside a region.

Measuring when the tree stops paying

The price of all this is sensitivity to dimension. To see it, build the tree above over 20,000 points drawn uniformly from the unit cube, with leaf size 16, and count the nodes the best-first search opens for a single nearest neighbour, averaged over 50 random queries. Each tree has 4,095 nodes.

DimensionsNodes opened, averageShare of the tree
2130.3%
4230.6%
81393.4%
161,96748%

These figures come from one run with a fixed random seed and will vary a little with yours, but the shape is robust. In low dimensions the tree opens a handful of nodes. By 16 dimensions it opens half of them and, once the bookkeeping is counted, loses to a vectorised linear scan. The reason is geometric: in high dimensions the distance to the nearest neighbour approaches the distance to a typical point, so the ball around the query intersects most boxes and almost nothing can be pruned.

Real data lying near a low-dimensional surface is kinder than uniform noise, so measure on your own data against a brute-force baseline.

Failure modes

  • Infinite build on duplicates. A split that cannot separate identical points recurses forever. Stop when the widest spread is zero.
  • Wrong answers from a bad bound. Mixing squared and plain distances, or using cosine similarity with a Euclidean bound, silently prunes true neighbours. Test against brute force with random inputs, including duplicates and queries outside the data.
  • Unscaled features. One feature measured in thousands dominates the distance, and the tree splits only on it. Standardise first.
  • High-dimensional slowness that looks fine in tests. Tests on 1,000 points never reveal it. Log nodes opened per query in production.
  • Stale trees. A static tree does not see inserts. Rebuild periodically and scan a small buffer of recent points alongside it.

Trade-offs and alternatives

Use a k-d tree when dimensions are low, roughly up to ten or so for exact search on typical data, when you need exact answers, and when the data is mostly static. In Python, scipy.spatial.KDTree and scikit-learn's KDTree are mature; in C++, nanoflann is a widely used header-only library. Scikit-learn's BallTree bounds nodes with spheres instead of boxes and can handle metrics that box bounds do not support.

For embedding vectors with hundreds of dimensions, an exact tree will degrade to a scan, so move to approximate graph or quantisation indexes such as HNSW or the disk-resident DiskANN. Between the two regimes, the honest answer is to benchmark: a well-vectorised brute-force scan over a few hundred thousand points is often faster than people expect, and it has no build time or staleness. The traversal patterns used here, explicit stacks and priority queues instead of recursion, are the same ones described in BFS and DFS.

What to do next

  1. Measure your data's dimensionality and scale; standardise features and drop ones that do not matter to the distance.
  2. Copy the implementation above and test it against a brute-force oracle on random data, including duplicates, single points and queries far outside the data.
  3. Add a counter for leaves scanned and plot it against dimension on your real data.
  4. Benchmark against a vectorised brute-force scan and a library tree such as SciPy's on the same machine; keep the fastest that meets your accuracy needs.
  5. If you need region counts or sums, add per-node aggregates and use the containment shortcut.
  6. Decide a rebuild policy for changing data: interval, change threshold and a buffer of recent points.
  7. If nodes opened approaches half the tree, switch to an approximate index and measure recall.
Key takeaway: A k-d tree is correct as long as its discard test is a true lower bound, so engineer it around better bounds. Store a tight bounding box per node, search best-first and stop when the nearest unopened box is farther than the k-th best point, and count range queries by adding whole nodes that lie inside the query. Test against brute force, guard against duplicates, and measure nodes opened on your real data: pruning fades quickly with dimension, and past that point a vectorised scan or an approximate index wins.