Day 038-Q1 — Suffix Array + Wavelet Tree(区間k番目文字クエリ)

2026-05-21 赤色 Master / Phase 8+ ★★★★★★★★★ SA / Wavelet Tree

問題

長さ $N$ の文字列 $S$(英小文字)と $Q$ 個のクエリ。クエリ l r k は、Suffix Array (SA) の $[l-1, r-1]$ (0-indexed) の範囲の suffix の先頭文字の Multiset から辞書順 $k$ 番目の文字を求める。

制約

$1 \le N \le 2 \times 10^5$
$1 \le Q \le 2 \times 10^5$
$1 \le l \le r \le N$, $1 \le k \le r-l+1$
時間制限: 2sec / メモリ: 256MB

入出力例

入力例 1

5 3
abcba
1 5 2
2 4 1
1 3 3

出力例 1

a
a
b

SA = [4,3,0,2,1] → 先頭文字 = ['a','b','a','c','b']。クエリ1: [a,a,b,b,c] → 2番目='a'。クエリ3: [a,b,a] → ソート → [a,a,b] → 3番目='b'。

概念図: Wavelet Tree 構造

アルファベット $[0, 25]$ を完全二分木で分割。各ノードは「左子に行く要素の累積カウント」を保持。

ルート [0..25] cnt[]: 各位置で左(≤12)に行く累積数 [0..12] (a-m) 左子: ≤ 6 [13..25] (n-z) 右子: ≥ 13 [0..6] (a-g) 深さ log₂(26)=5 [7..12] (h-m) ... [13..19] (n-t) ... [20..25] (u-z) ... kth(l, r, k): 各レベルで left_in_range = cnt[r+1] - cnt[l] を計算 k ≤ left_in_range → 左子へ (インデックスを左子空間に写像) k > left_in_range → k -= left_in_range して右子へ

ヒント(段階的開示)

ヒント1: 方向性
Suffix Array で suffix を辞書順ソート済みに並べた配列を取得する。クエリは「ある範囲の要素 multiset の k番目」なので、Wavelet Tree が適合する。
ヒント2: アプローチ
SA の各インデックスに「その suffix の先頭文字(0-25)」を格納した列 $A$ を作り、Wavelet Tree を構築。クエリ l r kwavelet.kth(l-1, r-1, k)。構築 $O(N \log \sigma)$、クエリ $O(\log \sigma)$。
ヒント3: kth の写像計算
# kth(l, r, k) — 0-indexed [l,r], 1-indexed k
# 各レベルで:
cnt = self.nodes[node_id]  # 累積カウント配列
left_in_range = cnt[r+1] - cnt[l]
if k <= left_in_range:
    # 左子へ (インデックス写像)
    r_new = cnt[r+1] - 1
    l_new = cnt[l]
    (l, r) = (l_new, r_new)
    hi = mid; node_id = 2*node_id+1
else:
    k -= left_in_range
    # 右子へ
    right_offset = cnt[n]  # 左子に行った総数
    r_new = (r+1 - cnt[r+1]) + right_offset - 1
    l_new = (l - cnt[l]) + right_offset
    # ※ 右子の配列は左に行かなかった要素の列
    lo = mid+1; node_id = 2*node_id+2

模範解答 (Python)

import sys
from sys import stdin
input = stdin.readline

def build_sa(s):
    """O(N log N) Doubling Suffix Array"""
    n = len(s)
    sa = list(range(n))
    rank = [ord(c) for c in s]
    tmp = [0] * n
    k = 1
    while k < n:
        def cmp_key(i):
            r2 = rank[i + k] if i + k < n else -1
            return (rank[i], r2)
        sa.sort(key=cmp_key)
        tmp[sa[0]] = 0
        for i in range(1, n):
            tmp[sa[i]] = tmp[sa[i-1]]
            if cmp_key(sa[i]) != cmp_key(sa[i-1]):
                tmp[sa[i]] += 1
        rank = tmp[:]
        if rank[sa[-1]] == n - 1:
            break
        k *= 2
    return sa

class WaveletTree:
    def __init__(self, arr, lo_val, hi_val):
        self.n = len(arr)
        self.lo = lo_val
        self.hi = hi_val
        self.nodes = {}
        self._build(arr, lo_val, hi_val, 0)

    def _build(self, arr, lo, hi, node_id):
        if lo == hi or not arr:
            return
        mid = (lo + hi) // 2
        left_arr, right_arr = [], []
        cnt = [0] * (len(arr) + 1)
        for i, v in enumerate(arr):
            if v <= mid:
                cnt[i+1] = cnt[i] + 1
                left_arr.append(v)
            else:
                cnt[i+1] = cnt[i]
                right_arr.append(v)
        self.nodes[node_id] = cnt
        self._build(left_arr, lo, mid, 2*node_id+1)
        self._build(right_arr, mid+1, hi, 2*node_id+2)

    def kth(self, l, r, k):
        lo, hi, node_id = self.lo, self.hi, 0
        while lo < hi:
            cnt = self.nodes[node_id]
            n_total = len(cnt) - 1
            left_in_range = cnt[r+1] - cnt[l]
            mid = (lo + hi) // 2
            if k <= left_in_range:
                r = cnt[r+1] - 1
                l = cnt[l]
                hi = mid
                node_id = 2*node_id+1
            else:
                k -= left_in_range
                right_offset = cnt[n_total]
                r = (r+1 - cnt[r+1]) + right_offset - 1
                l = (l - cnt[l]) + right_offset
                lo = mid + 1
                node_id = 2*node_id+2
        return lo

def solve():
    N, Q = map(int, input().split())
    S = input().strip()
    sa = build_sa(S)
    A = [ord(S[sa[i]]) - ord('a') for i in range(N)]
    wt = WaveletTree(A, 0, 25)
    out = []
    for _ in range(Q):
        l, r, k = map(int, input().split())
        c = wt.kth(l-1, r-1, k)
        out.append(chr(c + ord('a')))
    print('\n'.join(out))

solve()

Step-by-Step 解説

1Suffix Array 構築 (Doubling)
各反復で suffix を (rank[i], rank[i+k]) のペアでソート。$O(N \log^2 N)$(安定ソートなら $O(N \log N)$)。
2先頭文字配列の生成
A[i] = ord(S[sa[i]]) - ord('a')。SA 順(辞書順)の先頭文字を 0-25 の整数で格納。
3Wavelet Tree 構築
アルファベット $[0, 25]$ を完全二分木で表現。各ノードに「左子に行く要素の prefix sum」を格納。$O(N \log 26) = O(N)$。
4kth クエリ
各レベルで left_in_range = cnt[r+1] - cnt[l] を計算。$k \le$ left_in_range なら左子へ(インデックスを写像)、そうでなければ右子へ。$O(\log 26) = O(1)$ に近い定数時間。
5文字への変換と出力
整数 $c$ を chr(c + ord('a')) で文字に戻して出力。

計算量

SA 構築: $O(N \log^2 N)$
Wavelet Tree 構築: $O(N \log \sigma)$ — $\sigma = 26$
kth クエリ: $O(\log \sigma)$ per query
合計: $O((N + Q) \log \sigma + N \log^2 N)$

よくあるミス

ミス原因正しい書き方
Wavelet Tree の区間写像ミス左子への写像で cnt のオフセットを誤るr_new = cnt[r+1]-1, l_new = cnt[l]
右子への写像で right_offset を忘れる右子の配列インデックスが左子の続きになるl_new = (l - cnt[l]) + right_offset
SA の 0/1-indexed 混在クエリ変換忘れ入力は 1-indexed → l-1, r-1
葉判定なしで無限ループwhile 条件漏れwhile lo < hi: で lo==hi で停止

次のステップ

  • 発展: SA + LCP + Wavelet Tree で「区間内 k 番目に小さい部分文字列」を返す
  • 応用: 区間内の最頻文字クエリ(Wavelet Tree の range_freq 操作)
  • 実装: SA-IS (O(N)) + 動的 Wavelet Tree(オンラインクエリ対応)

自己評価

自分の回答

気づき・メモ