Day 102-Q2 — 一般化サフィックスオートマトン(Generalized SAM via Trie)

2026-07-25 赤色 Master / Phase 8+ ★★★★★★★★★ Trieベース構築・複数文字列の相異なる部分文字列数え上げ

問題

$K$個の文字列$S_1,\dots,S_K$(小文字英字のみ)が与えられる。これらすべてに現れる部分文字列を集めた集合(和集合、重複は1回のみ)のサイズ、すなわち相異なる非空部分文字列の総数を求めよ。

区切り文字で連結して1本のSAMを作ると「区切り文字をまたぐニセの部分文字列」を数えてしまい誤り。一般化サフィックスオートマトン(Generalized SAM)を使い、全文字列から作ったTrieをBFS順に走査しながらSAMの`extend`操作を呼ぶことで、複数文字列を正しくマージした1つのオートマトンを構築する。

入力形式

K
S_1
S_2
:
S_K

制約

$1 \le K \le 10^5$
各$S_i$は小文字英字の非空文字列
$\sum|S_i| \le 2\times10^5$

入出力例

入力例1

2
ab
ba

出力例1

4

"ab"→{a,b,ab}、"ba"→{b,a,ba}。和集合は{a,b,ab,ba}の4個。

入力例2

3
abc
bc
c

出力例2

6

"abc"→{a,b,c,ab,bc,abc}、"bc"→{b,c,bc}、"c"→{c}。和集合は{a,b,c,ab,bc,abc}の6個。

概念図

Trie構築 → BFS順にSAMへextend("ab","ba"の例) a b ab ba a b b a BFS順: 根→(a,b)→(ab,ba)。各辺でSAM.extend(親のSAM状態, 文字)を呼ぶ answer = Σ(len[v] - len[link[v]]) for v≠root = 4

ヒント(段階的開示)

ヒント1: 方向性
全文字列を`#`のような区切り文字で連結して普通のSAMを1回構築する方法は、区切り文字をまたぐ「偽の部分文字列」まで数えてしまうため誤り。文字列ごとに独立してSAMを作りPythonの`set`で和集合を取る方法は正しいが$O((\sum|S_i|)^2)$になり大きい制約では間に合わない。
ヒント2: アプローチ
全文字列からTrieを構築し、根をSAMの初期状態に対応づける。Trieを根から幅優先(BFS)でたどりながら、各Trie辺$(親,子,c)$に対して「親に対応するSAM状態」から文字$c$で`extend`を呼び、結果を「子に対応するSAM状態」として記録する。DFSでなくBFSを使うことが、`extend`内部の`last`不変条件を保つために重要。
ヒント3: 誘導(コード骨格)
from collections import deque
trie_children = [dict()]
for s in strings:
    cur = 0
    for c in s:
        if c not in trie_children[cur]:
            trie_children.append(dict())
            trie_children[cur][c] = len(trie_children) - 1
        cur = trie_children[cur][c]

trie_to_sam = [0] * len(trie_children)
q = deque((0, child, c) for c, child in trie_children[0].items())
while q:
    parent_trie, node_trie, ch = q.popleft()
    sam_node = sam.extend(trie_to_sam[parent_trie], ch)
    trie_to_sam[node_trie] = sam_node
    for c2, child2 in trie_children[node_trie].items():
        q.append((node_trie, child2, c2))
# answer = sum(len[v]-len[link[v]] for v in 1..sz-1)

`extend`自体は通常のSAM構築とほぼ同じだが、「すでにその文字への遷移が存在する場合」の分岐(clone判定)を先頭に追加する点が異なる。

模範解答 (Python)

import sys
from collections import deque


class GeneralizedSAM:
    def __init__(self, maxn):
        self.nxt = [dict() for _ in range(maxn)]
        self.link = [-1] * maxn
        self.length = [0] * maxn
        self.sz = 1  # 0番 = 根(空文字列に対応)

    def extend(self, last, c):
        nxt, link, length = self.nxt, self.link, self.length

        if c in nxt[last]:
            q = nxt[last][c]
            if length[q] == length[last] + 1:
                return q
            clone = self.sz; self.sz += 1
            length[clone] = length[last] + 1
            link[clone] = link[q]
            nxt[clone] = dict(nxt[q])
            p = last
            while p != -1 and nxt[p].get(c) == q:
                nxt[p][c] = clone
                p = link[p]
            link[q] = clone
            return clone

        cur = self.sz; self.sz += 1
        length[cur] = length[last] + 1
        p = last
        while p != -1 and c not in nxt[p]:
            nxt[p][c] = cur
            p = link[p]
        if p == -1:
            link[cur] = 0
        else:
            q = nxt[p][c]
            if length[q] == length[p] + 1:
                link[cur] = q
            else:
                clone = self.sz; self.sz += 1
                length[clone] = length[p] + 1
                link[clone] = link[q]
                nxt[clone] = dict(nxt[q])
                pp = p
                while pp != -1 and nxt[pp].get(c) == q:
                    nxt[pp][c] = clone
                    pp = link[pp]
                link[q] = clone
                link[cur] = clone
        return cur

    def distinct_substrings(self):
        total = 0
        for v in range(1, self.sz):
            total += self.length[v] - self.length[self.link[v]]
        return total


def solve():
    data = sys.stdin.read().split()
    K = int(data[0])
    strings = data[1:1 + K]

    trie_children = [dict()]
    for s in strings:
        cur = 0
        for c in s:
            if c not in trie_children[cur]:
                trie_children.append(dict())
                trie_children[cur][c] = len(trie_children) - 1
            cur = trie_children[cur][c]

    total_len = sum(len(s) for s in strings)
    sam = GeneralizedSAM(total_len * 2 + 5)

    trie_to_sam = [0] * len(trie_children)
    q = deque((0, child, c) for c, child in trie_children[0].items())
    while q:
        parent_trie, node_trie, ch = q.popleft()
        sam_node = sam.extend(trie_to_sam[parent_trie], ch)
        trie_to_sam[node_trie] = sam_node
        for c2, child2 in trie_children[node_trie].items():
            q.append((node_trie, child2, c2))

    print(sam.distinct_substrings())


solve()
計算量: $O(\sum|S_i|)$(Trie構築 + BFS + SAM extendはいずれも状態数・遷移数に線形)。

Step-by-Step 解説

1Trieの構築
全$K$本の文字列を1つのTrieにまとめる。共通prefixは自然に共有され、偽の部分文字列を生まない。
2BFS順でSAMに`extend`を適用
根からBFSでTrieをたどり、各辺$(親,c)$ごとに`extend(last,c)`を呼ぶ。BFSであることが不変条件維持の鍵。
3`extend`内の「既存遷移」分岐
複数文字列がTrie上でprefixを共有する場合「その遷移は既に存在している」ケースが頻発するため、先頭で既存遷移チェックとclone処理を行う。
4相異なる部分文字列の数え上げ
$\sum_{v\ne root}(\text{len}(v)-\text{len}(\text{link}(v)))$が、統合オートマトンの表す相異なる部分文字列の総数。

よくあるミス

ミス原因正しい書き方
区切り文字で連結して普通のSAMを1回構築区切り文字をまたぐ偽の部分文字列まで数えてしまうTrie+BFS+`extend`による正式な一般化SAM構築を行う
TrieをDFS順で処理する`extend`内部の`last`不変条件が崩れclone処理が誤ったノードを参照する必ずBFS(幅優先)でTrieを走査する
`extend`で「既存遷移がある場合」の分岐を書き忘れる単一文字列版のコードをそのまま流用する典型的な移植バグ先頭に`if c in nxt[last]:`のclone判定を追加する
ノード配列のサイズを$\sum|S_i|$のみで確保clone状態の分だけ余分にノードが必要(最大概ね2倍)`2\sum|S_i|+5`程度の余裕を持って確保する

次のステップ

  • 発展: 一般化SAMで「$K$本すべてに共通する最長部分文字列」を求める(各状態がどの文字列由来かをビットマスクで伝播させるDP)
  • 発展: 一般化SAM上のトポロジカル順(`length`のカウンティングソート)で各文字列ごとの出現回数(endpos集合サイズ)を求める
  • 次回予告: 下限付き最大流の実行可能性判定(Lower-Bound Feasible Flow・仮想ソース/シンク変換)

自己評価

自分の回答

気づき・メモ