Sulba
000 / 100

Subtree of Another Tree

EasyTime O(m + n)Space O(m + n)LeetCode 572 ↗

The problem

Given the roots of two binary trees, root and subRoot, return true if root contains a subtree — some node together with everything below it — identical to subRoot.

Examples

01
Input
root = [3, 4, 5, 1, 2], subRoot = [4, 1, 2]
Output
true
02
Input
root = [3, 4, 5, 1, 2, null, null, null, null, 0], subRoot = [4, 1, 2]
Output
false

Constraints

  • 1 ≤ nodes in root ≤ 2000
  • 1 ≤ nodes in subRoot ≤ 1000
  • −10⁴ ≤ Node.val ≤ 10⁴

The idea

Comparing subRoot against every node of root works but costs m × n. Instead turn both trees into text and search the text.

Write a tree in preorder (node, then left, then right), with # for every missing child and ^ before every value, so ^1 can never be mistaken for the end of ^21. With the gaps written in, the text fixes the tree exactly, and a subtree is precisely a run of its tree’s text. The Knuth–Morris–Pratt search finds a run in linear time: when a match fails part-way, a precomputed table says how much of what was matched can still be reused, so the text is never re-read.

Time
O(m + n) — writing both trees and one linear search
Space
O(m + n) — the two texts

Solution · every language run against every case

class Solution:    def isSubtree(self, root: Optional[TreeNode], subRoot: Optional[TreeNode]) -> bool:        # Write each tree in preorder, "^" before every value and "#" for every gap.        # A subtree is then exactly a run of the big tree's text.        def write(node, out):            if not node:                out.append("#")                return            out.append("^" + str(node.val))            write(node.left, out)            write(node.right, out)         text, pat = [], []        write(root, text)        write(subRoot, pat)        text, pat = "".join(text), "".join(pat)        # Knuth-Morris-Pratt: fail[i] is the longest proper prefix of pat[:i+1] that is also its suffix.        fail = [0] * len(pat)        k = 0        for i in range(1, len(pat)):            while k and pat[i] != pat[k]:                k = fail[k - 1]            if pat[i] == pat[k]:                k += 1            fail[i] = k        k = 0        for c in text:            while k and c != pat[k]:                k = fail[k - 1]            if c == pat[k]:                k += 1            if k == len(pat):                return True        return False