Day 100-Q3 — DC3法/Skew Algorithm(Suffix Arrayの線形時間構築)

2026-07-23 赤色 Master / Phase 8+ ★★★★★★★★★ DC3 / Skew Algorithm

問題

文字列 $S$(長さ $N$、英小文字のみ)の Suffix Array(接尾辞配列)を構築せよ。すべての接尾辞 $S[i:]$ を辞書式順序で昇順に並べたときの開始位置 $i$ の列を求める。

今回は $O(N)$ を達成する DC3法(Difference Cover mod 3、通称 Skew Algorithm) を使う。全接尾辞を $i \bmod 3$ で2群($B_0$: $i\equiv0$、$B_{12}$: $i\equiv1,2$)に分け、$B_{12}$ を3文字ずつまとめて再帰的に順位付けし、その順位を使って $B_0$ と統合(マージ)する分割統治法である。

入力形式

S

制約

$1 \le N \le 2\times10^5$
$S$ は英小文字のみ

入出力例

入力例1

banana

出力例1

5 3 1 0 4 2

a(5) < ana(3) < anana(1) < banana(0) < na(4) < nana(2) の辞書式順。

概念図

DC3: 1/3の要素だけ再帰、残り2/3をマージ index i: 0 1 2 3 4 5 6 B0 (i%3=0): 0 3 6 B12(i%3=1,2): 1 2 4 5 ① 3文字組を基数ソート → ランク付け(重複あれば再帰) ② B12のSuffix Array確定 ③ B0 を (S[i], rank(i+1)) のペアで基数ソート ④ 2ポインタでマージ → 最終Suffix Array

ヒント(段階的開示)

ヒント1: 方向性
SA-ISは「S/L型」分類から誘導ソートしたが、DC3は「文字列を3文字ずつのブロックに圧縮できれば、短い文字列の接尾辞ソートに帰着できる」という発想。$T(N)=T(2N/3)+O(N)$ の再帰式が$O(N)$に解けることを利用する。
ヒント2: アプローチ

1. $B_1\cup B_2$の各接尾辞の先頭3文字を1つの記号とみなし基数ソート・ランク付け。全部相異なればそのまま順位、重複があれば再帰で真の順位を確定。

2. $B_0$は$(S[i],\text{rank}(i+1))$のペアで基数ソート。

3. ソート済み$B_{12}$と$B_0$を先頭から比較しながらマージする。

ヒント3: 誘導(コード骨格)
def radix_pass(a, b, r, off, n, K):
    c = [0] * (K + 1)
    for i in range(n):
        c[r[off + a[i]]] += 1
    total = 0
    for i in range(K + 1):
        t = c[i]; c[i] = total; total += t
    for i in range(n):
        key = r[off + a[i]]
        b[c[key]] = a[i]
        c[key] += 1

文字列末尾には番兵として 0 を3つ余分に付ける。マージのstart位置は t = n0 - n1($N\bmod3=1$の境界調整)。

模範解答 (Python)

import sys


def radix_pass(a, b, r, off, n, K):
    c = [0] * (K + 1)
    for i in range(n):
        c[r[off + a[i]]] += 1
    total = 0
    for i in range(K + 1):
        t = c[i]
        c[i] = total
        total += t
    for i in range(n):
        key = r[off + a[i]]
        b[c[key]] = a[i]
        c[key] += 1


def leq2(a1, a2, b1, b2):
    return a1 < b1 or (a1 == b1 and a2 <= b2)


def leq3(a1, a2, a3, b1, b2, b3):
    return a1 < b1 or (a1 == b1 and leq2(a2, a3, b2, b3))


def dc3(s, n, K):
    if n == 0:
        return []
    if n == 1:
        return [0]
    n0 = (n + 2) // 3
    n1 = (n + 1) // 3
    n2 = n // 3
    n02 = n0 + n2

    s12 = [0] * (n02 + 3)
    SA12 = [0] * (n02 + 3)

    j = 0
    for i in range(n + (n0 - n1)):
        if i % 3 != 0:
            s12[j] = i
            j += 1

    radix_pass(s12[:n02], SA12, s, 2, n02, K)
    tmp = SA12[:n02]
    radix_pass(tmp, s12, s, 1, n02, K)
    tmp = s12[:n02]
    radix_pass(tmp, SA12, s, 0, n02, K)

    name = 0
    c0 = c1 = c2 = -1
    for i in range(n02):
        idx = SA12[i]
        v0, v1, v2 = s[idx], s[idx + 1], s[idx + 2]
        if v0 != c0 or v1 != c1 or v2 != c2:
            name += 1
            c0, c1, c2 = v0, v1, v2
        if idx % 3 == 1:
            s12[idx // 3] = name
        else:
            s12[idx // 3 + n0] = name

    if name < n02:
        sa12_rec = dc3(s12[:n02] + [0, 0, 0], n02, name)
        for i in range(n02):
            s12[sa12_rec[i]] = i + 1
        SA12 = sa12_rec
    else:
        SA12 = [0] * n02
        for i in range(n02):
            SA12[s12[i] - 1] = i

    s0 = [0] * n0
    SA0 = [0] * n0
    j = 0
    for i in range(n02):
        if SA12[i] < n0:
            s0[j] = 3 * SA12[i]
            j += 1
    radix_pass(s0[:n0], SA0, s, 0, n0, K)

    def get_i(t):
        return SA12[t] * 3 + 1 if SA12[t] < n0 else (SA12[t] - n0) * 3 + 2

    SA = [0] * n
    p = 0
    t = n0 - n1
    k = 0
    while k < n:
        i = get_i(t)
        jidx = SA0[p]
        if SA12[t] < n0:
            cond = leq2(s[i], s12[SA12[t] + n0], s[jidx], s12[jidx // 3])
        else:
            cond = leq3(s[i], s[i + 1], s12[SA12[t] - n0 + 1],
                        s[jidx], s[jidx + 1], s12[jidx // 3 + n0])
        if cond:
            SA[k] = i
            t += 1
            k += 1
            if t == n02:
                while p < n0:
                    SA[k] = SA0[p]
                    p += 1
                    k += 1
                break
        else:
            SA[k] = jidx
            p += 1
            k += 1
            if p == n0:
                while t < n02:
                    SA[k] = get_i(t)
                    t += 1
                    k += 1
                break
    return SA


def suffix_array(text):
    n = len(text)
    mapped = [ord(c) - ord('a') + 1 for c in text]
    s = mapped + [0, 0, 0]
    return dc3(s, n, 26)


def solve():
    text = input().strip()
    sa = suffix_array(text)
    print(*sa)


solve()
計算量: $T(N)=T(\lceil 2N/3\rceil)+O(N)$ より $O(N)$。Python定数倍のため実測は$N\le2\times10^5$で1秒未満。

Step-by-Step 解説

1群分けとインデックス列挙
$i\bmod3\ne0$の$i$を昇順に列挙。範囲をn+(n0-n1)とし境界調整を吸収する。
23文字組の基数ソートとランク付け
3回のLSD基数ソートで辞書式ソートし、隣接比較でランクを振る。全て相異なれば最終順位として使える。
3再帰呼び出し
重複があればランク列を新しい文字列として再帰適用。入力サイズは常に$2/3$以下になり全体$O(N)$。
4$B_0$のソートとマージ
$(S[i],\text{rank}(i+1))$で基数ソートし、$B_{12}$と$O(1)$比較でマージする。
5実装上の注意
理論$O(N)$だがPythonの定数倍は大きい。$N$が更に大きい場合はPyPy実行を検討。

よくあるミス

ミス原因正しい書き方
マージ開始位置t0にする$n_0$と$n_1$の差(0か1)だけ先頭に仮想エントリが挟まる構造を見落とすt = n0 - n1から開始する
末尾のパディングを2つ以下にする3文字組の参照が文字列の実長を超える番兵0を3つ追加する
基数ソートを不安定なsorted(key=...)で代用DC3は3回のLSD基数ソートの安定性に依存カウンティングソートベースの安定な実装にする
再帰呼び出し時にランク配列の末尾パディングを忘れるs12[:n02]だけ渡すと範囲外参照になるs12[:n02] + [0,0,0]として渡す

次のステップ

  • 発展: Kasaiのアルゴリズムで LCP配列 を $O(N)$ 導出し、Sparse Tableと組み合わせて任意の2接尾辞間LCPクエリを $O(1)$ にする
  • 次回予告: 杜教篩(Du's Sieve・乗法的関数の前綴和高速計算)

自己評価

自分の回答

気づき・メモ