Skip to content
AITroveRead. Build. Understand.
Make this comfortable

Ball trees: prune exact nearest-depot search with radius bounds

Last updated: 4 Oct 20269 min read
tutorial
IntermediateBy AITrove Editorial

A ball tree groups static points in recursively nested metric balls. Each node stores a center and radius that enclose all points beneath it. For a query point, the triangle inequality gives a lower distance bound: distance to the center minus the radius, floored at zero. If that bound exceeds the best depot distance found so far, the entire node can be skipped without losing an exact answer. This example uses Euclidean coordinates, chooses a split axis by largest spread, divides at the median, and checks every point only inside small leaves. Children are visited in lower-bound order so an early close depot can tighten pruning. This is a static nearest-neighbor structure; it neither inserts moved depots nor guarantees that every query avoids a full scan.

Operational case

For six depots at (0,0), (4,1), (9,0), (4,7), (-3,5), and (10,8), the nearest to (5,2) is D-47 at distance about 1.414. The nearest to (8,1) is D-61 at the same rounded distance; near (0,4), D-29 wins at about 3.162. The tree compares real Euclidean distances and breaks equal-distance ties by depot ID. It reports three decimal places only for display. A node's center need not be one of the stored depots, and its radius is the farthest enclosed point from that center. Pruning on a center distance alone would be unsound.

Working Python program

python
from math import dist


class DepotBall:
    def __init__(self, depots, leaf_size):
        dimensions = len(depots[0][1])
        self.center = tuple(sum(point[axis] for _, point in depots) / len(depots)
                            for axis in range(dimensions))
        self.radius = max(dist(self.center, point) for _, point in depots)
        self.left = None
        self.right = None
        self.depots = None
        if len(depots) <= leaf_size:
            self.depots = depots
        else:
            spread = [max(point[axis] for _, point in depots) - min(point[axis] for _, point in depots)
                      for axis in range(dimensions)]
            axis = max(range(dimensions), key=lambda position: spread[position])
            ordered = sorted(depots, key=lambda depot: depot[1][axis])
            middle = len(ordered) // 2
            self.left = DepotBall(ordered[:middle], leaf_size)
            self.right = DepotBall(ordered[middle:], leaf_size)

    def lower_bound(self, query):
        return max(0.0, dist(self.center, query) - self.radius)


class DepotBallIndex:
    def __init__(self, depots, leaf_size=2):
        if not depots or leaf_size < 1:
            raise ValueError("nonempty depots and positive leaf size required")
        dimensions = len(depots[0][1])
        if dimensions < 1 or any(len(point) != dimensions for _, point in depots):
            raise ValueError("consistent positive dimension required")
        self.root = DepotBall(list(depots), leaf_size)
        self.dimensions = dimensions

    def nearest(self, query):
        if len(query) != self.dimensions:
            raise ValueError("query dimension mismatch")
        best = None

        def visit(node):
            nonlocal best
            if best is not None and node.lower_bound(query) > best[0] + 1e-12:
                return
            if node.depots is not None:
                for depot_id, point in node.depots:
                    candidate = (dist(query, point), depot_id)
                    if best is None or candidate < best:
                        best = candidate
                return
            children = sorted((node.left, node.right), key=lambda child: child.lower_bound(query))
            for child in children:
                visit(child)

        visit(self.root)
        return best[1], round(best[0], 3)


depot_points = [
    ("D-19", (0, 0)), ("D-47", (4, 1)), ("D-61", (9, 0)),
    ("D-83", (4, 7)), ("D-29", (-3, 5)), ("D-37", (10, 8)),
]
depot_index = DepotBallIndex(depot_points)
print(depot_index.nearest((5, 2)), depot_index.nearest((8, 1)))
print(depot_index.nearest((0, 4)))

Output

Output
('D-47', 1.414) ('D-61', 1.414)
('D-29', 3.162)

Time, space, and tradeoff

At each balanced level this Python builder sorts each subset along a selected axis, giving O(N log squared N) construction time in the worst case and O(N) retained nodes and points. A query may inspect all N points when balls overlap or bounds are weak; there is no universal logarithmic search guarantee. Each visited node computes a center distance, and recursion uses O(log N) depth for the median split. Keeping an exact leaf scan matters even after bound pruning. This model uses floating-point geometry with a small pruning tolerance; application code needing strict numerical guarantees must set a precision policy and test boundary cases.

Common Mistakes

  • Do not prune from center distance without subtracting the enclosing radius.
  • Do not claim every ball-tree nearest query is logarithmic.
  • Do not return a center as though it were a stored depot.
  • Do not silently mix coordinate dimensions in one tree.

Connected lessons

Compare this operation boundary with Successor disjoint sets: skip permanently retired slots, Persistent range-distinct counts: keep only the latest position active, Editable substring fingerprints: join hashes in a segment tree, Sliding medians: expire heap entries by event identity, then complete the audit project and decision quiz.

data structures
range-query-structures
Storage details