第 73 章 Trie 字典树¶
配套例题:BISHI93 【模板】Trie 字典树 来源:S4
模板.docx「字符串 → 字典树(trie 树)」(含指针版与数组版两份 C++ 模板) 前置:07-字典、70-字符串处理、46-位运算
Trie(字典树 / 前缀树)是把一堆字符串按公共前缀合并成一棵树的数据结构。 它解决的是哈希表解决不了的那一类问题:
| 问题 | set / dict 能做吗 |
Trie |
|---|---|---|
| 某个串在不在集合里 | ✅ \(O(1)\),比 Trie 快 | ✅ \(O(\|s\|)\) |
| 有多少个串以 \(t\) 为前缀 | ❌ 得枚举所有串 | ✅ \(O(\|t\|)\) |
| 集合里与 \(x\) 异或最大的数 | ❌ | ✅ 01-Trie,\(O(\log V)\) |
| 集合里字典序最小/第 \(k\) 小的串 | ❌ 得排序 | ✅ 树上二分 |
| 多个模式串同时在文本里匹配 | ❌ | ✅ AC 自动机(Aho-Corasick 自动机)= Trie + KMP |
一句话判据:问「精确匹配」用哈希表,问「前缀 / 逐位贪心」用 Trie。 只要题面里出现「前缀」「异或最大」「按位从高到低」这些词,就该想 Trie。
这一章的重点不是 Trie 的原理(它非常简单),而是 Python 下三种实现的取舍—— 同一个算法,三种写法在 CPython 里的内存能差 30 倍,这直接决定题能不能过。
73.1 结构与用途¶
一棵 Trie 的边上是字符,节点代表一个前缀(根代表空前缀):
插入 wangzai / waylon / wangle 之后:
(root)
│w
● pre=3
┌──┴──┐
a│ │
● pre=3 (前缀 "wa")
┌──┴──┐
n│ │y
● ● pre=2 / pre=1
│g │l
● ● pre=2 (前缀 "wang")
┌─┴─┐ ⋮
z│ │l
● ● pre=1 / pre=1
⋮ ⋮
每个节点维护两个计数就够应付绝大多数题:
| 字段 | 含义 | 更新时机 |
|---|---|---|
pre[v] |
经过节点 \(v\) 的插入次数 = 以该前缀开头的串数 | 插入时沿途每个节点 \(+1\) |
end[v] |
终止于节点 \(v\) 的插入次数 = 恰等于该串的个数 | 插入时只在最后一个节点 \(+1\) |
这两个计数千万别混。「有多少个串以
wang为前缀」问的是pre, 「集合里有几个wang」问的是end。BISHI93 问的是pre—— 之所以答案能一步查出来,是因为 「以 \(t\) 为前缀的串」恰好就是「插入时经过了 \(t\) 对应节点的串」。
三个基本操作,复杂度全是 \(O(|s|)\),与集合里有多少个串无关:
| 操作 | 做法 |
|---|---|
插入 s |
从根沿字符走,没路就建节点;沿途 pre += 1,末端 end += 1 |
查询 s 是否存在 |
沿字符走,走不通返回否;走通了看 end > 0 |
前缀计数 s |
沿字符走,走不通返回 0;走通了返回 pre |
删除 s |
沿途 pre -= 1,末端 end -= 1(节点不必真删) |
空间上界:节点数 \(\le\) 所有串的总长度 \(+1\)。BISHI93 的总长 \(10^6\), 所以最坏有 \(10^6+1\) 个节点——这个上界是后面所有内存讨论的基准。
⚠️ 一个必须先破除的错误直觉:既然要查前缀,为什么不直接 「把每个串的所有前缀切出来丢进
Counter」?# [片段] ❌ 看着聪明,实则致命 from collections import Counter cnt = Counter() for s in words: for i in range(1, len(s) + 1): cnt[s[:i]] += 1 # 切片 s[:i] 是 O(i) 的拷贝总长虽然只有 \(10^6\),但它可能是一个长度 \(10^6\) 的串。 切出它的全部前缀要拷贝 \(1+2+\cdots+10^6 \approx 5\times10^{11}\) 个字符—— 时间和内存双爆。Trie 存在的意义就是让公共前缀只存一份。
73.2 Python 下的三种实现¶
三种实现的算法完全一样,差别只在「一个节点的 \(C\) 条出边存在哪」。
实现一:嵌套 dict¶
每个节点就是一个 dict,键是字符,值是子节点(也是 dict)。
def trie_nested_build(words):
"""嵌套 dict 版 Trie。'#' 存 pre 计数,'$' 存 end 计数。
最好写,但每个节点一个 dict 对象,内存代价最大。
"""
root = {} # 根代表空前缀
for w in words:
node = root # 每个串都从根出发
for ch in w:
node = node.setdefault(ch, {}) # ★ 「查 + 不存在就新建」一次调用完成
node['#'] = node.get('#', 0) + 1 # 沿途累加 = 经过该前缀的串数(pre)
node['$'] = node.get('$', 0) + 1 # 只在末端累加 = 恰好等于该串的个数(end)
return root
def trie_nested_count_prefix(root, t):
"""有多少个串以 t 为前缀。"""
node = root
for ch in t:
node = node.get(ch)
if node is None: # 走不通 -> 没有任何串以 t 为前缀
return 0 # 断路返回,不能继续往下走
return node['#'] # '#' 槽里存的就是 pre 计数
setdefault 是这个写法的灵魂:「查 + 不存在就插入」一次调用完成,
比 if ch not in node: node[ch] = {} 少一次哈希。
| 优点 | 缺点 |
|---|---|
| 代码最短,五行就能写出来 | 每个节点一个 dict 对象 |
| 天然只占用出现过的边 | 空 dict 光对象头就 64 字节,\(10^6\) 个节点 = 64 MB 起,加上装入元素后的哈希表实际远超 100 MB |
| 字符集任意大(Unicode 也行) | 每层都要一次属性/方法调用,常数最大 |
调试时 print(root) 直接看见结构 |
把计数塞进同一个 dict('#'/'$')很容易和真实字符撞车 |
⚠️
'#'这种标记键的撞车风险是真实存在的。BISHI93 保证只有字母, 所以安全;但题目若允许任意可见字符,就必须换成不可能出现的键 (比如整数0——dict的键可以混类型)或者干脆用下面两种实现。
适用规模:节点数 \(\le 10^5\)。写小题、写面试题、写 \(n \le 1000\) 的暴力对照,用它。
实现二:单个扁平 dict(Python 的最优解)¶
只开一个 dict,把「节点编号」和「字符」编码进同一个整数键:
def trie_flat_build(words, C=128):
"""单个扁平 dict 版 Trie(BISHI93 的做法)。
words 传 bytes 列表;key = node * C + 字符字节值 -> 子节点编号。
C 取 128 覆盖全部 ASCII,且是 2 的幂(乘法可换成移位)。
返回 (child, pre, end)。
"""
child = {} # 全局只有这一个 dict,键是「节点 + 字符」的编码
pre = [0] # pre[0] 是根,不会被查到
end = [0] # len(pre) 同时充当「下一个可用节点编号」计数器
get = child.get # ★ 绑成局部名,省掉每次属性查找
for w in words:
cur = 0 # 0 号节点是根,每个串都从这里出发
for b in w: # 迭代 bytes 得到的是 int
k = cur * C + b # 与扁平数组同一套下标编码,只是槽不预先分配
nxt = get(k, -1) # 一次调用完成「查表 + 缺省」,-1 表示没这条边
if nxt < 0:
nxt = len(pre) # 新节点编号 = 当前节点总数
pre.append(0)
end.append(0)
child[k] = nxt
cur = nxt
pre[cur] += 1 # 沿途每个节点 +1 = 经过该前缀的串数
end[cur] += 1 # 只有末端 +1 = 恰好等于该串的个数
return child, pre, end
def trie_flat_count_prefix(child, pre, t, C=128):
"""有多少个串以 t 为前缀。走不通立刻返回 0。"""
get = child.get
cur = 0
for b in t:
cur = get(cur * C + b, -1) # 用同一套编码往下走一步
if cur < 0:
return 0 # 断路:这条边不存在,就没有串以 t 为前缀
return pre[cur] # 走通了,答案就是终点节点的 pre
为什么这是 Python 下的最优解,四条理由:
| 理由 | 说明 |
|---|---|
只有一个 dict 对象 |
内存 = 「实际存在的边数」× 每条 entry 的开销,没有任何按 \(C\) 分配的浪费 |
| 整数键的哈希是恒等映射 | CPython 里 hash(小整数) == 整数本身,查表比字符串键快得多 |
get(k, -1) 一次调用完成「查 + 缺省」 |
比 if k in child: ... else: ... 少一次哈希 |
计数用两个 list |
整数下标访问,比往 dict 里塞 '#' 快 |
C为什么取 128 而不是 26 或 52: 直接用原始字节值当转移字符,就不必写b - 97这种减法, 也不必为「大小写混排」做映射(BISHI93 区分大小写,字母有 52 种但 ASCII 码分布在 65–90 和 97–122 两段,减法映射反而更麻烦)。 128 是 2 的幂,cur * 128 + b在 CPython 里和cur << 7 | b几乎等速, 而且保证b < 128时编码无歧义。
适用规模:节点数 \(\le 2\times10^6\)。BISHI93 的 \(10^6\) 就落在这里。
实现三:扁平数组(C++ 的最优解,Python 的陷阱)¶
S4 模板.docx 的第二份模板 int next[MAX][26] 就是这个:给每个节点固定分配 \(C\) 个槽。
def trie_array_build(words, C=26, base=97):
"""扁平数组版 Trie:son[node * C + c] = 子节点编号,0 表示不存在。
对应 S4 的 int next[MAX][26]。0 号节点是根,所以「0 = 不存在」不会歧义。
查询是纯下标访问,理论上最快;但内存是 节点数 * C,Python 下极易爆。
"""
son = [0] * C # 先给 0 号节点(根)开一段,占 son[0..C-1]
pre = [0] # len(pre) 就是当前节点总数,兼作编号分配器
end = [0]
for w in words:
cur = 0 # 每个串都从根出发
for b in w:
# ★ 二维 next[node][c] 压成一维:每个节点独占连续 C 个槽,
# node 号的那一段从 node*C 开始,第 c 条边落在 node*C + c
k = cur * C + (b - base) # b - base 把字符映射到 0..C-1
nxt = son[k]
if nxt == 0: # 0 号是根,根不可能是谁的孩子,故 0 = 不存在
nxt = len(pre) # 新节点编号 = 已分配的节点数
son[k] = nxt
son.extend([0] * C) # ★ 为新节点开一整段,哪怕只用一条边
pre.append(0)
end.append(0)
cur = nxt
pre[cur] += 1 # 沿途每个节点 +1 = 经过该前缀的串数
end[cur] += 1 # 只有末端 +1 = 恰好等于该串的个数
return son, pre, end
在 C++ 里这是最优解,在 Python 里几乎总是错的选择:
C++ int next[1e6][26] |
Python son = [0] * (26 * 1e6) |
|
|---|---|---|
| 每个槽 | 4 字节 int |
8 字节指针 + 指向的 int 对象 |
| 总内存 | \(10^6 \times 26 \times 4 = 104\) MB | 指针数组 \(2.6\times10^7 \times 8 = 208\) MB,再加 \(10^6\) 个非缓存小整数对象(每个 28 字节)约 28 MB |
| 初始化 | memset,一瞬间 |
[0] * 2.6e7 要 0.2 s 以上,extend \(10^6\) 次还有扩容抖动 |
| 结论 | ✅ 首选 | ❌ 512 MB 的空间限制下直接 MLE |
反过来说,它什么时候在 Python 里可用:\(C\) 很小、节点数也很小的时候。 典型就是 01-Trie(\(C = 2\))——\(10^5\) 个数 × 30 位 = \(3\times10^6\) 个节点,
ch数组 \(6\times10^6\) 个槽,约 50 MB,可以接受。见 73.4。
选型表¶
| 实现 | 内存(\(N\) 个节点,字符集 \(C\)) | 常数 | 可行节点数(512 MB) | 什么时候用 |
|---|---|---|---|---|
嵌套 dict |
\(O(N)\),但每节点 64 B 起 | 大 | \(\sim10^5\) | 手写快、小数据、面试 |
单个扁平 dict |
\(O(\text{边数})\),最省 | 小 | \(\sim2\times10^6\) | 竞赛默认选择 |
| 扁平数组 | \(O(N \cdot C)\) | 最小 | \(C=26\) 时 \(\sim2\times10^6\) 槽 ⇒ \(N \sim 8\times10^4\) | 只在 \(C \le 4\)(01-Trie)时用 |
决策一句话: 字符集是字母 → 单个扁平
dict;字符集是 \(\{0,1\}\) → 扁平数组。 嵌套dict只用来写十行内的小工具。
73.3 插入、查询与前缀统计¶
把实现二包成一个可复用的模板(含 end 计数与删除):
class Trie:
"""字符串 Trie,单扁平 dict 实现。字符集为字节值 0..127。
- insert(w) 插入一个 bytes/str
- count_prefix(t) 有多少个已插入的串以 t 为前缀
- count_exact(t) 有多少个已插入的串恰好等于 t
- erase(w) 删除一次 w(必须保证 w 已被插入过)
"""
__slots__ = ("child", "pre", "end")
def __init__(self):
self.child = {} # key = node * 128 + 字节值 -> 子节点编号
self.pre = [0] # 0 号节点是根
self.end = [0]
def insert(self, w):
if isinstance(w, str):
w = w.encode() # 统一成 bytes,迭代出来才是 int
child, pre, end = self.child, self.pre, self.end
get = child.get # 循环里高频调用,先绑成局部名
cur = 0 # 从根出发
for b in w:
k = cur * 128 + b # 每个节点虚拟占 128 个槽,实际只存用到的边
nxt = get(k, -1)
if nxt < 0: # 没这条边就新建一个节点
nxt = len(pre) # 新编号 = 已有节点数
pre.append(0)
end.append(0)
child[k] = nxt
cur = nxt
pre[cur] += 1 # 沿途 +1:经过该前缀的串数
end[cur] += 1 # 末端 +1:恰好等于 w 的串数
def _walk(self, t):
"""沿 t 走,返回终点节点编号;走不通返回 -1。"""
if isinstance(t, str):
t = t.encode()
get = self.child.get
cur = 0
for b in t:
cur = get(cur * 128 + b, -1)
if cur < 0:
return -1 # 断路退出,绝不能拿 -1 去索引 pre/end
return cur
def count_prefix(self, t):
v = self._walk(t)
return 0 if v < 0 else self.pre[v]
def count_exact(self, t):
v = self._walk(t)
return 0 if v < 0 else self.end[v]
def erase(self, w):
"""删除一次 w。只减计数,不回收节点——回收得不偿失。"""
if isinstance(w, str):
w = w.encode()
get = self.child.get
cur = 0
path = [] # 记下沿途节点,确认能删之后再统一减
for b in w:
cur = get(cur * 128 + b, -1)
if cur < 0:
return False # 不存在,什么都不做
path.append(cur)
if self.end[cur] == 0:
return False # 路走得通,但没有串「终止」在这里
for v in path:
self.pre[v] -= 1 # 沿途撤销一次「经过」
self.end[cur] -= 1 # 末端撤销一次「终止」
return True
四个高频派生查询,都不需要改结构:
| 查询 | 做法 |
|---|---|
| 集合里所有串的最长公共前缀 | 从根往下走,只要当前节点 pre == n 且只有一个孩子就继续 |
| 给定 \(s\),集合里是 \(s\) 前缀的串有几个 | 沿 \(s\) 走,把沿途所有 end 加起来 |
| 按字典序输出所有串 | 对 Trie 做 DFS,孩子按字符升序访问——Trie 的 DFS 序天然是字典序 |
| 字典序第 \(k\) 小的串 | 从根往下,每一步用子树里的串数(pre 或 end 之和)做二分 |
⚠️ 递归深度:Trie 的高度等于最长串的长度。 遍历 Trie 时若写递归 DFS,串长 \(10^6\) 就意味着 \(10^6\) 层递归—— 必爆 C 栈且无任何报错。Trie 上的遍历一律写显式栈, 理由与 90.5 完全相同。
73.4 01-Trie 与最大异或对¶
问题¶
给定 \(n\) 个非负整数,求 \(\max_{i \ne j} (a_i \oplus a_j)\)。
\(n \le 10^5\) 时暴力两两枚举是 \(5\times10^9\) 次异或,Python 下毫无希望 (即使 C++ 也要好几秒)。
为什么 Trie 能救¶
异或的关键性质是按位独立(见 46-位运算):
要让 \(a \oplus b\) 最大,就要让最高位尽可能是 1; 最高位定了之后,次高位再尽可能是 1……高位的一个 1 顶得上后面所有位。
于是:把每个数按二进制从高位到低位当成一个长度固定的「串」插进 Trie (字符集 \(C = 2\),这就是 01-Trie)。查询 \(x\) 时从根往下, 每一位都优先走与 \(x\) 当前位相反的分支——走得通,这一位的异或就是 1。
插入 3(011) 与 5(101) 之后,查询 x = 4(100):
root
0/ \1
● ●
1/ 0\
● ●
1/ \1
● ●
(3) (5)
x 的最高位是 1 -> 优先走 0 分支,通 -> 该位异或得 1
次高位是 0 -> 优先走 1 分支,通 -> 该位异或得 1
最低位是 0 -> 优先走 1 分支,通 -> 该位异或得 1
=> 4 ^ 3 = 7 ✓
模板¶
\(C = 2\),正是扁平数组唯一划算的场合:
def max_xor_pair(nums, bits=30):
"""最大异或对:从 nums 里选两个数使异或最大。O(n * bits)。
ch[v * 2 + b] = 节点 v 沿 b 走到的孩子,0 表示不存在(0 号是根)。
bits 必须覆盖数据上界:a < 2^30 取 30,a <= 1e9 取 30,a <= 1e18 取 60。
"""
if len(nums) < 2:
return 0
ch = [0, 0] # 根(0 号节点)的两条边,占 ch[0] 与 ch[1]
for x in nums: # ---- 建树 ----
cur = 0
for k in range(bits - 1, -1, -1): # ★ 必须从高位往低位
b = (x >> k) & 1 # 取出 x 的第 k 位
# 字符集只有 {0,1},每个节点独占 2 个槽:v 号的边在 ch[v*2] 和 ch[v*2+1]
nxt = ch[cur * 2 + b]
if nxt == 0: # 根是 0 号,不会是谁的孩子,故 0 = 不存在
nxt = len(ch) >> 1 # 新节点编号 = 已有节点数 = 槽数的一半
ch[cur * 2 + b] = nxt
ch.append(0) # 补出新节点自己的两个槽
ch.append(0)
cur = nxt
best = 0
for x in nums: # ---- 逐个查询 ----
cur = 0
val = 0 # x 与最优配对的异或值,逐位拼出来
for k in range(bits - 1, -1, -1):
b = (x >> k) & 1
opp = ch[cur * 2 + (b ^ 1)] # b ^ 1 就是 b 的相反位
if opp: # 能走相反位 -> 这一位异或得 1
val |= 1 << k # 高位的一个 1 顶得上后面所有位,能拿就拿
cur = opp
else: # 只能走同位 -> 这一位是 0
cur = ch[cur * 2 + b] # 每个数都插满 bits 位,同位分支一定存在
if val > best:
best = val
return best
三个必查的点:
| 点 | 说明 |
|---|---|
bits 要够 |
\(a \le 10^9 < 2^{30}\) 取 30;写小了高位直接丢失,答案偏小 |
| 必须从高位往低位 | 反了就变成「让低位尽量是 1」,完全错 |
| 同一批数先全插入再全查询 | 这样 \(i = j\) 也会被枚举到,但 \(x \oplus x = 0\) 不影响最大值;若题目要求 \(i < j\),改成「边插边查」(第一个数只插不查) |
Python 的更优路线:用 set 代替 Trie¶
01-Trie 的主循环是 \(n \times \text{bits}\) 次纯 Python 层迭代。 \(n = 10^5\)、\(bits = 30\) 就是 \(3\times10^6\) 次建树 + \(3\times10^6\) 次查询, 每次还有两三次列表下标访问——实测 4–6 秒,在 2 秒时限下过不了。
同一个贪心可以改写成「按位试填 + 集合查表」,把内层循环全部下沉到 C 层:
def max_xor_pair_set(nums, bits=30):
"""最大异或对的无 Trie 写法:按位试填答案。O(bits * n),但全在 C 层。
原理:设答案的高位已确定为 res,试着让第 k 位也为 1(记 cand = res | 1<<k)。
只看每个数的高 (bits-k) 位(记作前缀 p),则
存在一对数满足高位异或 == cand <=> 存在 p 使 p ^ cand 也是前缀。
集合推导式与 `in` 判断都是 C 层操作,比逐位走 Trie 快一个数量级。
"""
res = 0 # 已经坐实的答案高位
mask = 0 # 只保留「已考察过的高位」
for k in range(bits - 1, -1, -1): # 与 01-Trie 一样,必须从高位往低位定
mask |= 1 << k # 每轮把考察范围往低位放宽一位
prefixes = {x & mask for x in nums} # ★ 集合推导,C 层一次扫完
cand = res | (1 << k) # 试着让第 k 位也取 1
for p in prefixes: # a ^ b == cand <=> b == a ^ cand
if (p ^ cand) in prefixes:
res = cand # 确实存在这样一对,第 k 位坐实为 1
break # 本位已定,进入下一位
return res
| 01-Trie | 按位试填 + set |
|
|---|---|---|
| 复杂度 | \(O(n \log V)\) | \(O(n \log V)\)(同阶) |
| Python 层迭代次数 | \(2n\log V \approx 6\times10^6\) | \(\log V\) 次集合推导(每次 C 层扫 \(n\) 个) |
| 实测(\(n = 10^5\),30 位) | 4–6 s ❌ | 0.2–0.4 s ✅ |
| 能否回答「与某个固定 \(x\) 异或最大」 | ✅ 直接查 | ❌ 只求全局最大 |
| 能否支持在线插入 / 删除 | ✅ | ❌ 每次都要重扫 |
| 能否做「异或第 \(k\) 大」「区间版」 | ✅ | ❌ |
决策: - 只问「全局最大异或对」→ 按位试填 +
set,这是 Python 下唯一稳过的写法; - 要「对每个 \(x\) 分别求最优配对」「支持插入删除」「异或第 \(k\) 大」「限定下标区间」 → 只能 01-Trie,并且要认命:Python 下 \(n \cdot \log V \le 2\times10^6\) 才现实。
01-Trie 的其它套路¶
在每个 01-Trie 节点上再存一个 cnt(子树里有多少个数),就能做:
| 需求 | 做法 |
|---|---|
| 与 \(x\) 异或第 \(k\) 大 | 每一位看「相反分支的 cnt」够不够 \(k\),够就走相反分支,否则 \(k\) 减掉后走同位 |
| 有多少对数的异或 \(< L\) | 沿 \(L\) 的二进制走,在 \(L\) 的每个 1 位处把「异或该位为 0」的整棵子树的 cnt 累加 |
| 异或值在 \([l, r]\) 的对数 | 拆成 \(f(r+1) - f(l)\),\(f\) 用上一行 |
| 区间 \([l, r]\) 内与 \(x\) 异或最大 | 可持久化 01-Trie(每插入一个数建一个新版本,共享未改动的子树) |
| 最小异或生成树 | Borůvka + 01-Trie(见 92.7) |
Python 的现实提醒:上表越往下,Python 层的循环次数越多。 可持久化 01-Trie 的节点数是 \(n \log V \approx 3\times10^6\), 光建树就要好几秒——这类题在 CPython 下基本属于「写出来练手,比赛里放弃」。
73.5 Trie 的其它用途¶
| 用途 | 说明 |
|---|---|
| AC 自动机(Aho-Corasick 自动机) | Trie + 失配指针(KMP 的 \(\pi\) 数组在树上的推广),一次扫文本匹配全部模式串。见 71.8 |
| 字典序遍历 | Trie 的 DFS 序就是字典序,常用来「按字典序输出所有满足条件的串」 |
| 后缀 Trie / 后缀自动机 | 把所有后缀插进 Trie,可解本质不同子串数等问题;后缀 Trie 是 \(O(n^2)\) 的,实用的是后缀自动机(超纲) |
| 数对 / 路径的异或 | 树上路径异或 = 根到两点的前缀异或的异或,于是「树上最大异或路径」= 01-Trie 求最大异或对 |
| 压缩 Trie(基数树) | 把只有一个孩子的链压成一条边,节点数降到 \(O(\text{串数})\);Python 里写起来麻烦,很少用 |
一个漂亮的转化:「树上两点路径异或最大」的做法是 先一遍 DFS 求出 $d[v] = $ 根到 \(v\) 的边权异或, 则路径 \(u \to v\) 的异或 \(= d[u] \oplus d[v]\)(LCA 上方的部分被异或两次抵消)。 于是问题变成「\(n\) 个数里选两个异或最大」。前缀异或 + 抵消是异或题的通用手法, 见 42-前缀和与差分 与 94-树上算法。
73.6 例题¶
BISHI93 【模板】Trie 字典树(中等)¶
给定 \(n\) 个模式串 \(s_1..s_n\) 与 \(q\) 次查询,第 \(i\) 次给一个文本串 \(t_i\), 统计有多少个模式串以 \(t_i\) 为前缀。 \(1 \le n, q \le 10^5\),所有输入串的总长度 \(\le 10^6\),全部由大小写字母组成, 区分大小写。时限:C/C++ 1 秒,其他语言 2 秒;空间:其他语言 512 MB。 题面见 BISHI93 原题(牛客)。 题解见
solutions/BISHI93.py(已通过官方样例验证)。
标准 Trie 模板题:插入时沿途 pre += 1,查询时沿 \(t\) 走到底输出 pre。
真正要做的判断是选哪种实现——总长 \(10^6\) 意味着最坏 \(10^6\) 个节点,
73.2 的选型表直接给出答案:单个扁平 dict。
import sys
def main():
data = sys.stdin.buffer.read().split() # 串里没有空格,按 token 切正好一行一个
n, q = int(data[0]), int(data[1])
child = {} # key = node * 128 + 字节值 -> 子节点编号
cnt = [0] # cnt[v] = 经过节点 v 的模式串数;0 号是根
get = child.get # ★ 绑成局部名,2e6 次调用省下的属性查找很可观
p = 2
for _ in range(n): # ---- 插入 ----
s = data[p]; p += 1
cur = 0 # 每个模式串都从 0 号节点(根)出发
for b in s: # 迭代 bytes 得到 int,直接当转移字符
k = cur * 128 + b # 节点号乘字符集大小再加字符,编码唯一且可逆
nxt = get(k, -1) # 一次调用完成「查表 + 缺省」,-1 = 没这条边
if nxt < 0: # 没这条边就新建节点
nxt = len(cnt) # len(cnt) 就是已有节点数,直接拿来当新编号
cnt.append(0)
child[k] = nxt
cur = nxt
cnt[cur] += 1 # ★ 沿途每个节点都 +1,这就是「前缀计数」
out = []
for _ in range(q): # ---- 查询 ----
s = data[p]; p += 1
cur = 0
for b in s:
cur = get(cur * 128 + b, -1) # 用与插入完全相同的编码往下走
if cur < 0: # 走不通:没有任何模式串以 t 为前缀
break # 必须立刻断路,否则 cnt[-1] 会取到最后一个元素
out.append("0" if cur < 0 else str(cnt[cur])) # 走通了就输出终点的 cnt
sys.stdout.write("\n".join(out) + "\n") # 1e5 行一次写出,逐行 print 会拖垮时限
main()
复杂度 \(O(\sum|s| + \sum|t|) \le O(2\times10^6)\) 次转移,
每次转移是一次 dict.get。实测约 1.2–1.5 秒,2 秒时限能过但不宽裕。
五个坑:
- 区分大小写。样例第三问
Way的答案是 0,就是专门在测这个—— 任何lower()/ 统一映射都会答错。直接用原始字节值当转移字符最省心, 这也是取 \(C = 128\)(而不是 26 或 52)的原因; - 查询串可能比任何模式串都长,中途没有子节点必须立刻
break并输出 0。 写成「走完再判」会在cnt[-1]上取到最后一个元素,得到随机答案; cnt统计的是「经过」而不是「终止」。若写成只在末端+1, 查w会得到 0(没有模式串恰好等于w),而正确答案是 3;- 不能对每个前缀切片。总长 \(10^6\) 可能集中在一个串上, 切全部前缀是 \(5\times10^{11}\) 次字符拷贝(见 73.1 的警告框);
- IO 必须整块读。\(n + q = 2\times10^5\) 行,逐行
input()会额外花掉 0.3 秒以上, 在 1.2 秒的算法时间上是压垮时限的最后一根稻草。见 20-输入输出处理。
为什么不用嵌套
dict? \(10^6\) 个节点就是 \(10^6\) 个dict对象。 空dict的sys.getsizeof是 64 字节,装入一两个元素后会分配哈希表, 实测单节点 200 字节量级——总计 200 MB 以上,加上解释器开销逼近 512 MB 的上限, 而且每层多一次setdefault调用,时间上也过不了。为什么不用扁平数组? \(C = 52\)(或 128)× \(10^6\) 个节点 = \(5\times10^7\) 个列表槽,光指针就 400 MB,必 MLE。
本题给出的判据:Trie 在 Python 里能不能做,看的是「节点数 × 字符集」。 「节点数」由总串长决定,「字符集」由实现方式决定—— 单个扁平
dict把后者从 \(C\) 降成了「实际存在的边数」, 这就是它在 Python 下不可替代的原因。
73.7 本章速查¶
| 要点 | 结论 |
|---|---|
| Trie 是什么 | 把字符串按公共前缀合并成的树,边是字符,节点是前缀 |
| 什么时候用 Trie | 题面出现「前缀」「异或最大」「按位从高到低」 |
| 什么时候用哈希表 | 只问「精确匹配 / 存在性」——set 比 Trie 快 |
pre[v] vs end[v] |
经过 \(v\) 的串数 vs 终止于 \(v\) 的串数,别混 |
| 「以 \(t\) 为前缀的串数」 | 沿 \(t\) 走到底的 pre,\(O(\|t\|)\),与集合大小无关 |
| 节点数上界 | 所有串的总长度 \(+1\) |
❌ 切出所有前缀丢 Counter |
单个长 \(10^6\) 的串就是 \(5\times10^{11}\) 次拷贝 |
实现一:嵌套 dict |
最好写,setdefault 一行;每节点 64 B 起,\(\le 10^5\) 节点 |
实现二:单个扁平 dict |
竞赛默认:key = node*128 + 字节值,\(\le 2\times10^6\) 节点 |
| 实现三:扁平数组 | C++ 最优、Python 陷阱;\(C \cdot N\) 个指针,只在 \(C \le 4\) 时用 |
| 扁平 dict 的四个技巧 | 单个 dict / 整数键哈希恒等 / get(k, -1) / 计数用 list |
| \(C\) 取 128 | 直接用字节值,免映射,2 的幂,天然区分大小写 |
| 删除 | 只减 pre/end,不回收节点 |
| Trie 上遍历 | 必须显式栈——树高 = 最长串长,递归必爆栈 |
| Trie 的 DFS 序 | 天然是字典序(孩子按字符升序) |
| 01-Trie 原理 | 按位从高到低插入,查询时优先走相反分支 |
01-Trie 的 bits |
必须覆盖上界:\(10^9 \to 30\),\(10^{18} \to 60\)。写小了直接错 |
| 最大异或对(Python) | 按位试填 + set(\(0.2\)–\(0.4\) s),而不是 01-Trie(4–6 s) |
| 树上最大异或路径 | 前缀异或 \(d[u] \oplus d[v]\) ⇒ 退化成最大异或对 |
| AC 自动机 | Trie + 失配指针,多模式串匹配(71 章) |
| 规模 | Python 现实性 |
|---|---|
字符 Trie,总长 \(\le 10^6\)(扁平 dict) |
✅ 约 1.2–1.5 s |
字符 Trie,总长 \(\le 10^6\)(嵌套 dict) |
❌ 内存 200 MB+,时间也过不了 |
| 字符 Trie,总长 \(\le 10^6\)(扁平数组 \(C=52\)) | ❌ \(5\times10^7\) 个槽,必 MLE |
| 字符 Trie,总长 \(\ge 5\times10^6\) | ❌ 单纯的 \(10^7\) 次 dict.get 就要 5 s |
| 01-Trie,\(n\log V \le 2\times10^6\) | ⚠️ 勉强(\(n \le 6\times10^4\),30 位) |
| 01-Trie,\(n = 10^5\),30 位 | ❌ 4–6 s;改用按位试填 + set |
| 可持久化 01-Trie,\(n = 10^5\) | ❌ \(3\times10^6\) 个节点,建树就超时 |
| 看到什么 → 想到什么 |
|---|
| 「有多少个串以 … 为前缀」 |
| 「最大异或对」 |
| 「与 \(x\) 异或第 \(k\) 大」 |
| 「树上路径异或最大」 |
| 「多个模式串同时匹配文本」 |
| 「按字典序输出所有 …」 |
| 「所有串的最长公共前缀」 |