~/dynamic-programming/tree-dp
Dynamic programming on trees
Solve a tree children first: each node turns its children's answers into its own and hands one small summary up to its parent.
A post-order pass. Each node combines its children's summaries, updates a global answer if it needs one, and returns a summary to its parent.
The input is a tree and a node's answer depends only on its subtrees: path sums, diameters, picking non-adjacent nodes, counts below each node.
O(n)
O(n)
You’ll recognise it when
- The input is a tree: a
TreeNode, a parent array, orn - 1edges. You want one best value or count over it, or one per node. - A node’s answer depends only on its own value and its children’s answers: subtree sizes, heights, sums, “is this subtree balanced?”.
- The question is about paths that can start and end anywhere: the longest path, or the path with the largest sum.
- You pick or label nodes under a rule about neighbours: never a parent and its child together, every node watched by a neighbour, neighbours in different colours.
- Trying every start node, or every subset of nodes, is far too slow.
It looks like a plain post-order traversal, and it is one. The DP part is deciding what each call returns.
The idea
Think of a company that wants one number from the whole organisation, like the largest team. Nobody reads the full org chart. Each manager asks each direct report for a short summary, adds their own part, and passes one summary up to their own boss. Every person is asked once, and the answer arrives at the top.
That’s tree DP. A subtree is a smaller problem of the same kind, and different subtrees never overlap. So you solve the children first and build each node’s answer from theirs. The DP table has one entry per node, and post-order is the order that fills it.
The hard part is choosing what to report. Ask: what does my parent need from me? Often it isn’t the final answer. Take the largest sum of any path, where a path may start and end at any nodes and values can be negative:
- Every path has one highest node, where it bends. There, the path is the node plus the best path going straight down into its left side, plus the same on its right. Either side can be left out if it would lower the sum.
- The parent can only use a path that goes straight down from the node. A path that has already bent at the node can’t also go up to the parent: the node would have three neighbours on it.
So each call returns down, the best path that starts at the node and goes straight down, and on the side it checks bend against a global best. The tree traversal article shows this shape for the diameter. Here it is for path sums:
def max_path_sum(root):
best = root.val
def down(node): # best path straight down
nonlocal best
if node is None:
return 0
left = max(down(node.left), 0) # drop a losing side
right = max(down(node.right), 0)
# the path that bends here
best = max(best, node.val + left + right)
# what the parent can extend
return node.val + max(left, right)
down(root)
return best
How it works
The code in the tabs does the same work without recursion, on a tree with some negative values, so some sides are worth dropping. Scroll through the steps and the graphic follows along. You can also turn on quiz me to predict what each node hands up, or edit the tree.
- Order. List the nodes so that every child comes before its parent. Pop nodes off a stack, push their children, and read the list backwards. The small numbers show the order.
- Leaves first. -2 is a leaf. Its missing children count as 0, so the only path that bends at it is -2 itself.
- Hand up
down. -2 reportsdown = -2: the best path that starts at -2 and goes straight down. Its parent will read it from the table. - Drop a losing side. At 5, the left child hands up -2. Adding it would only lower the sum, so 5 counts it as
max(-2, 0) = 0. - Two helpful sides. At 6, the left side hands up 5, and the right side hands up 1: the path -3 → 4.
- Bend. The best path that turns at 6 takes both sides:
bend = 6 + 5 + 1 = 12. - Update
best. 12 beats the best so far, 5. The green edges mark the best path found so far. - Return one side. 6 hands its parent
down = 6 + max(5, 1) = 11, not 12. A path coming from the parent can only carry on down one side of 6. - Answer. After the root,
best = 12: the path 5 → 6 → -3 → 4. It bends at 6, not at the root, and it keeps -3 because the 4 below more than pays for it.
Why it’s right: every path in a tree has exactly one highest node. At that node, the path is the node plus at most one downward path into each child, so its best possible sum is that node’s bend, as long as each child’s down is right. down is right by induction from the leaves: a downward path is either the node alone, or the node plus a downward path from one child. Every node’s bend is checked against best, so the best path can’t be missed.
class TreeNode:
def __init__(self, val, left=None, right=None):
self.val = val
self.left = left
self.right = right
def max_path_sum(root):
"""Largest sum of any non-empty path, in one bottom-up pass."""
# Parents before children; read backwards, children come first.
order, stack = [], [root]
while stack:
node = stack.pop()
order.append(node)
for child in (node.left, node.right):
if child:
stack.append(child)
down = {None: 0} # best path from a node straight down; an empty child gives 0
best = root.val # one node on its own is a path
for node in reversed(order):
left = max(down[node.left], 0) # @step:children a losing branch counts as 0
right = max(down[node.right], 0)
bend = node.val + left + right # @step:bend the best path turning here
best = max(best, bend)
down[node] = node.val + max(left, right) # @step:return the parent extends one side only
return best#include <algorithm>
#include <unordered_map>
#include <vector>
using namespace std;
struct TreeNode {
int val;
TreeNode* left = nullptr;
TreeNode* right = nullptr;
explicit TreeNode(int v) : val(v) {}
};
// Largest sum of any non-empty path, in one bottom-up pass.
long long maxPathSum(TreeNode* root) {
// Parents before children; read backwards, children come first.
vector<TreeNode*> order, stack = {root};
while (!stack.empty()) {
TreeNode* node = stack.back();
stack.pop_back();
order.push_back(node);
for (TreeNode* child : {node->left, node->right})
if (child) stack.push_back(child);
}
// Best path from a node straight down; an empty child gives 0.
unordered_map<TreeNode*, long long> down = {{nullptr, 0}};
long long best = root->val; // one node on its own is a path
for (auto it = order.rbegin(); it != order.rend(); ++it) {
TreeNode* node = *it;
long long left = max(down[node->left], 0LL); // @step:children a losing branch counts as 0
long long right = max(down[node->right], 0LL);
long long bend = node->val + left + right; // @step:bend the best path turning here
best = max(best, bend);
down[node] = node->val + max(left, right); // @step:return the parent extends one side only
}
return best;
}import java.util.ArrayDeque;
import java.util.ArrayList;
import java.util.Deque;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
class TreeNode {
int val;
TreeNode left, right;
TreeNode(int val) { this.val = val; }
}
class MaxPathSum {
// Largest sum of any non-empty path, in one bottom-up pass.
static long maxPathSum(TreeNode root) {
// Parents before children; read backwards, children come first.
List<TreeNode> order = new ArrayList<>();
Deque<TreeNode> stack = new ArrayDeque<>();
stack.push(root);
while (!stack.isEmpty()) {
TreeNode node = stack.pop();
order.add(node);
if (node.left != null) stack.push(node.left);
if (node.right != null) stack.push(node.right);
}
// Best path from a node straight down; an empty child gives 0.
Map<TreeNode, Long> down = new HashMap<>();
down.put(null, 0L);
long best = root.val; // one node on its own is a path
for (int i = order.size() - 1; i >= 0; i--) {
TreeNode node = order.get(i);
long left = Math.max(down.get(node.left), 0); // @step:children a losing branch counts as 0
long right = Math.max(down.get(node.right), 0);
long bend = node.val + left + right; // @step:bend the best path turning here
best = Math.max(best, bend);
down.put(node, node.val + Math.max(left, right)); // @step:return the parent extends one side only
}
return best;
}
}The dictionary down holds what the recursive calls would have returned. It is the DP table: one entry per node, each filled once, after its children. The recursive version is shorter, but a tree shaped like a long chain makes it one call deeper per node, and Python stops at about 1,000 calls. sys.setrecursionlimit lets it go deeper, but a very deep recursion can then crash the interpreter itself. C++ and Java go much deeper, but a chain of a million nodes can overflow their stacks too. The explicit order has no depth limit.
Sometimes one number isn’t enough. In DP basics, you rob houses on a street, never two neighbours. Put the houses on a tree, where a parent and its child are neighbours. Whether a node can be robbed now depends on its parent’s choice, which the node can’t see. So each call returns two numbers, one for each choice, and lets the parent decide:
def rob(node):
"""(best if node is robbed, best if node is skipped)"""
if node is None:
return 0, 0
l_rob, l_skip = rob(node.left)
r_rob, r_skip = rob(node.right)
take = node.val + l_skip + r_skip # children skipped
skip = max(l_rob, l_skip) + max(r_rob, r_skip) # free
return take, skip
The answer is max(rob(root)). The same trick covers any rule about neighbours: return one number per state the parent cares about, like one count per colour, or “has a camera, watched, not watched”.
Why it’s O(n)
Each node is visited once. In a binary tree it does a constant amount of work with its children’s summaries. In a general tree it does work in proportion to its children, and all the children add up to n - 1. So the whole pass is O(n).
A node’s answer covers its whole subtree, but it only reads its children’s numbers, never the subtree itself. That’s where the saving comes from. The obvious alternative, starting a search from every node, repeats almost everything:
| Approach | Time | Extra space |
|---|---|---|
| Search from every node, best sum to any end | O(n²) | O(n) |
| Tree DP, recursive | O(n) | O(h) call stack: O(log n) balanced, O(n) a chain |
| Tree DP, explicit order | O(n) | O(n) for the order and the table |
Common mistakes
Handing up the bend
The parent needs a path it can extend, and a path that bends at the node can’t be extended. Returning the bend makes the parent count a fork, like 4 → 2 ← 5 plus 2 → 1 as one path.
return node.val + left + right # ✗ a fork
return node.val + max(left, right) # ✓ one side
Keeping a side that loses
A child’s down can be negative. Adding it can only hurt, so count it as 0 and leave that side out. Clamp the children, not the node itself: the node is always on its own path.
left = down[node.left] # ✗ drags it down
left = max(down[node.left], 0) # ✓ drop it
Starting best at 0
A path must have at least one node. If every value is negative, the answer is the largest single value, not 0 for an empty path.
best = 0 # ✗ 0 for [-3, -1, -8]
best = root.val # ✓ -1 for [-3, -1, -8]
Asking the children twice
A balanced check that calls height() at every node walks each subtree again from every ancestor: O(n²) on a chain. Return the height and the verdict from the same call, with -1 for “not balanced”.
abs(height(n.left) - height(n.right)) <= 1 # ✗ walks again
if l < 0 or r < 0 or abs(l - r) > 1: # ✓ one pass
return -1
return 1 + max(l, r)
Variations
- Diameter. The same shape with lengths instead of sums: return the height, and check the two tallest child heights added together. In a tree with any number of children, keep the two largest values from the children.
- Balanced check. Return the height, or -1 as soon as any subtree is out of balance. A node is balanced if both children are and their heights differ by at most 1.
- Trees given as edges or parents. Root the tree at any node. Get a parents-first order with a stack or BFS, recording each node’s parent, then walk it backwards and add each node’s summary into its parent’s:
size[parent[v]] += size[v]counts everyone below each node. - Several states per node. Return one number per state: robbed or skipped, one count per colour, or “has a camera, watched, not watched”. The parent combines the children’s states under its rule.
- Every node as the root. When the question asks for an answer with each node as the root, solve it for one root, then run a second pass from the top that moves the root across each edge and fixes up the answer in O(1). That’s rerooting: two passes, O(n), instead of one pass per root.
Check yourself
5 quick questions. Pick an answer to see why it's right or wrong.
-
1
Each node of a binary tree holds a positive number. You pick a set of nodes with the largest total, but you may never pick both a node and its parent. Which approach is correct?
Whether a node can be picked depends only on whether its parent is, so each subtree needs two answers: picked and skipped. Combining them children first gives the best total in O(n). The level-based ideas assume a whole level is picked or skipped together, but the best set can pick a node on the left of a level and skip its neighbour on the right. Greedy fails when one big node blocks two children that add up to more.
-
2
What does this print?
class TreeNode:def __init__(self, val, left=None, right=None):self.val, self.left, self.right = val, left, rightdef rob(node):if node is None:return (0, 0) # (rob this node, skip it)a, b = rob(node.left), rob(node.right)take = node.val + a[1] + b[1] # children must be skippedskip = max(a) + max(b) # children are free to choosereturn (take, skip)root = TreeNode(2, TreeNode(6, TreeNode(3), TreeNode(4)), TreeNode(5, None, TreeNode(7)))print(rob(root))The leaves give (3, 0), (4, 0) and (7, 0). Node 6: take 6, skip 3 + 4 = 7, so (6, 7). Node 5: take 5, skip 7, so (5, 7). The root: take = 2 + 7 + 7 = 16, using the grandchildren 3, 4 and 7; skip = max(6, 7) + max(5, 7) = 14.
(16, 11)is what you get if skipping a node forced both children to be robbed: skip must let each child choose its better option. -
3
This maximum path sum has a bug. What does it print? (The correct answer for this tree is 11.)
class TreeNode:def __init__(self, val, left=None, right=None):self.val, self.left, self.right = val, left, rightdef max_path(root):best = float("-inf")def go(node):nonlocal bestif node is None:return 0left = max(go(node.left), 0)right = max(go(node.right), 0)best = max(best, node.val + left + right)return node.val + left + rightgo(root)return bestroot = TreeNode(1, TreeNode(2, TreeNode(4), TreeNode(5)), TreeNode(3))print(max_path(root))Node 2 returns its bend, 2 + 4 + 5 = 11, instead of its best downward path, 2 + max(4, 5) = 7. The root then adds 1 + 11 + 3 = 15, which counts a “path” that forks at 2: node 2 would need three neighbours. The fix is
return node.val + max(left, right), and the answer becomes 11 (4 → 2 → 5, or 5 → 2 → 1 → 3). -
4
Your recursive tree DP passes small tests but fails in Python on a tree of 100,000 nodes that forms one long chain. What is the most reliable fix?
A chain makes the recursion 100,000 calls deep, and Python stops at about 1,000. Reading a parents-first order backwards gives children before parents with no recursion at all.
@cacheremoves repeated work, not depth. Raising the limit lets Python go deeper than its own C stack can hold, so it can crash the interpreter outright. Pre-order would be just as deep, and it visits a node before its children have answered. -
5
For the recursive maximum path sum on a tree with n nodes and height h, what are the time and the extra space?
Each node is visited once and does O(1) work with its children’s summaries, so O(n) time. A node’s answer covers its subtree, but it reads only the children’s numbers, never the subtree itself. The call stack holds one path from the root: O(h), which is O(log n) for a balanced tree but O(n) for a chain. It’s never O(1): the recursion or the explicit order takes space.
Practice problems
Solve these right here, in Python, C++ or Java. Tests run as you go.