import Base # Balanced fork trees for bounded, data-parallel Bend work. # # A fork tree mirrors a list of n items as a binary tree of depth ceil(log2(n)) # with n leaves. Because the split is balanced, every root-to-leaf path performs # the same amount of work, so a runtime that schedules lanes by tree depth # finishes all items together instead of stranding one slow lane. # # Bend's `!` device scheduler and its bounded machine stack both want this # shape: depth is O(log n) instead of O(n), and siblings are independent so a # runtime may evaluate them in parallel. # # The source list must be in a fixed deterministic order. `build` consumes # exactly the first `n` items and returns the unconsumed tail, so a caller walks # a list with a running cursor without ever indexing or splitting it # destructively. `flatten` and `fold` recover that order. # # Elements must be `Data`: the unconsumed tail is returned as a `+List`, so a # caller can hand the same tail to the next `build` and keep it. That reuse is # the point -- it is what makes the walk allocation-free and order-preserving. # # Generic parameters follow the Base convention: the element quantity and kind # are erased leading arguments, exactly as in `List.length`. type Tree is Data: Empty{} Leaf{value: A} Fork{left: Tree, right: Tree} # A tree plus the unconsumed tail of the source list. `build` consumes a # prefix, so the tail is what the next `build` call will see. type Built is Data: Built{tree: Tree, rest: +List} def built_tree(a, -A: Data, built: Built) -> Tree: Built{tree, rest} = built tree def built_rest(a, -A: Data, built: Built) -> +List: Built{tree, rest} = built rest # Fuel `build` needs for a tree over n elements. # # The recurrence rounds *up*: an odd count splits into `n/2` and `n - n/2`, # which differ by one, and the larger half is what bounds the recursion. Using # a plain halving count here under-counts by one, and the symptom is silent -- # the deepest leaves come back as `Empty`, so a caller evaluates fewer items # than it budgeted and nothing reports an error. # # `n == 1` still costs one unit: at zero fuel `build` matches its out-of-fuel # case before it ever reaches the leaf case. def depth(fuel: Nat, n: U32) -> Nat: match fuel n: case 0n _: 0n case _ 0: 0n case _ 1: 1n case 1n+f _: 1n + depth(f, ((n + 1) / 2 : U32)) # Build a balanced tree over the first `count` elements of `items`. # # `fuel` must be at least `depth(count)`. It is a parameter rather than a # constant so callers can state the bound explicitly and the termination checker # can accept the recursion. `balanced` computes it for you. def build(a, -A: Data, +fuel: Nat, +count: U32, +items: +List) -> Built: match fuel count items: # Out of fuel, nothing to build, or nothing left: no tree, keep the input. case 0n _ _: Built{Empty{}, items} case _ 0 _: Built{Empty{}, items} case _ _ Nil{}: Built{Empty{}, Nil{}} # One element is a leaf; the tail is handed back untouched. case _ 1 h <> t: Built{Leaf{h}, t} case 1n+n _ _: +half = (count / 2 : U32) +left = build(a, A, n, half, items) +right = build(a, A, n, (count - half : U32), built_rest(a, A, left)) Built{Fork{built_tree(a, A, left), built_tree(a, A, right)}, built_rest(a, A, right)} # Build over exactly `count` elements using precisely the fuel that needs. def balanced(a, -A: Data, +count: U32, +items: +List) -> Built: build(a, A, depth(U32.to_nat(count) + 1n, count), count, items) # Number of leaves. For a tree built by `build` this equals the count asked for. def size(a, -A: Data, tree: Tree) -> U32: match tree: case Empty{}: 0 case Leaf{value}: 1 case Fork{left, right}: (size(a, A, left) + size(a, A, right) : U32) # Levels from this node to the deepest leaf. For a balanced tree this equals # `depth(count)`, the bound `build` relied on. def height(a, -A: Data, tree: Tree) -> Nat: match tree: case Empty{}: 0n case Leaf{value}: 0n case Fork{left, right}: 1n + Nat.max(height(a, A, left), height(a, A, right)) # In-order traversal. Because `build` splits a list prefix at its midpoint, # this returns the consumed prefix in its original order. Any fold that must # respect source order -- sample-index order, deterministic tie-breaking -- is # written against this and stays correct. def flatten(a, -A: Data, +tree: Tree) -> +List: match tree: case Empty{}: Nil{} case Leaf{value}: [value] case Fork{left, right}: List.append(&2, A, flatten(a, A, left), flatten(a, A, right)) # Left fold, in source order, over the whole tree. Following `List.foldl`, the # type parameters and the step are templates: the whole fold is inlined at # compile time, so the step costs nothing at runtime. def fold(~a: Quant, ~A: Kind(a), ~B: Type, ~f: A -> B -> B, tree: Tree, init: B) -> B: match tree: case Empty{}: init case Leaf{value}: f(value, init) case Fork{left, right}: fold(~a, ~A, ~B, ~f, right, fold(~a, ~A, ~B, ~f, left, init))