Height-Balanced Tree Check
Problem
A search index keeps its keys in a binary tree and rebuilds it whenever it gets lopsided. A tree is height-balanced when, at every node, the heights of the left and right subtrees differ by at most 1. Given the root of a binary tree, return True if it is height-balanced. The tree has up to 5,000 nodes, so computing the height again from every node is wasteful.
Examples
Input: tree = 3, left 9, right 20 (children 15 and 7)
Output: True
Input: tree = 1, left 2 (left child 3 with children 4 and 4, right child 3), right 2
Output: False
Why: at the root the left side is 3 levels deep and the right side only 1
Input: tree = empty
Output: True
Why: edge case, an empty tree has nothing out of balance
Hints
0 / 3
Calling a height function from every node visits the lower nodes again and again, which is quadratic on a tall tree.
A node needs its children's heights anyway to test itself, and its own height is one more than the larger of them. One bottom-up walk can compute both.
Return the height of each subtree from a post-order walk, or -1 as a signal that some subtree below is already unbalanced. Pass -1 straight up as soon as you see it.
Solution
A single post-order walk returns each subtree's height to its parent, so every node can check its own balance from numbers its children already worked out. When a node finds its two heights differ by more than one, it returns -1 instead of a height, and every ancestor passes that -1 straight up without doing more work. The tree is balanced exactly when the root does not return -1. Every node is visited once, so time is O(n), and space is O(h) for the recursion.
class T:
def __init__(self, val, left=None, right=None):
self.val, self.left, self.right = val, left, right
def is_balanced(root):
def height(node): # height, or -1 if unbalanced below
if node is None:
return 0
lh = height(node.left)
if lh == -1:
return -1
rh = height(node.right)
if rh == -1 or abs(lh - rh) > 1:
return -1
return 1 + max(lh, rh)
return height(root) != -1
print(is_balanced(T(3, T(9), T(20, T(15), T(7))))) # -> True
print(is_balanced(T(1, T(2, T(3, T(4), T(4)), T(3)), T(2)))) # -> False
print(is_balanced(None)) # -> TrueStuck on the idea rather than the code? Tree Depth and Balance covers it.