跳转至

第 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,把「节点编号」和「字符」编码进同一个整数键:

\[\texttt{key} = \texttt{node} \times C + \texttt{ch}\]
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\)的串 从根往下,每一步用子树里的串数(preend 之和)做二分

⚠️ 递归深度: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 秒时限能过但不宽裕

五个坑

  1. 区分大小写。样例第三问 Way 的答案是 0,就是专门在测这个—— 任何 lower() / 统一映射都会答错。直接用原始字节值当转移字符最省心, 这也是取 \(C = 128\)(而不是 26 或 52)的原因;
  2. 查询串可能比任何模式串都长,中途没有子节点必须立刻 break 并输出 0。 写成「走完再判」会在 cnt[-1] 上取到最后一个元素,得到随机答案;
  3. cnt 统计的是「经过」而不是「终止」。若写成只在末端 +1, 查 w 会得到 0(没有模式串恰好等于 w),而正确答案是 3;
  4. 不能对每个前缀切片。总长 \(10^6\) 可能集中在一个串上, 切全部前缀是 \(5\times10^{11}\) 次字符拷贝(见 73.1 的警告框);
  5. IO 必须整块读\(n + q = 2\times10^5\) 行,逐行 input() 会额外花掉 0.3 秒以上, 在 1.2 秒的算法时间上是压垮时限的最后一根稻草。见 20-输入输出处理

为什么不用嵌套 dict \(10^6\) 个节点就是 \(10^6\)dict 对象。 空 dictsys.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\) 大」
「树上路径异或最大」
「多个模式串同时匹配文本」
「按字典序输出所有 …」
「所有串的最长公共前缀」