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))