Day 015-Q1 — ウェーブレット木(Wavelet Tree)

2026-04-28 赤色 Master / Phase 8+ ★★★★★★★★★ k番目・以下個数

問題

長さ N の整数列 A に対し: kth l r k は A[l..r] の k 番目に小さい値、count l r v は v 以下の個数。

制約

$1 \le N, Q \le 2 \times 10^5$
$1 \le A_i \le 10^9$

入出力例

入力例 1

7 3
3 1 4 1 5 9 2
kth 1 5 3
count 2 6 4
kth 3 7 2

出力例 1

3
3
4

ヒント (段階的開示)

ヒント1: 方向性
値域をビット分解し二分木で管理。O(log V) で k番目・以下個数。
ヒント2: アプローチ
各ノードで「左子に行く要素数」を prefix sum で管理。座標圧縮で深さ O(log N)。
ヒント3: 誘導
k 番目: left_cnt 比較で左右どちらに降りるか決定。

模範解答 (Python)

import sys
from bisect import bisect_right
input = sys.stdin.readline

class WaveletTree:
    def __init__(self, arr, lo, hi):
        self.n = len(arr)
        self.lo = lo; self.hi = hi
        self.left = None; self.right = None
        if lo == hi:
            self.left_cnt = [0] * (self.n + 1)
            return
        mid = (lo + hi) // 2
        self.left_cnt = [0] * (self.n + 1)
        left_arr = []; right_arr = []
        for i, v in enumerate(arr):
            if v <= mid:
                left_arr.append(v)
                self.left_cnt[i+1] = self.left_cnt[i] + 1
            else:
                right_arr.append(v)
                self.left_cnt[i+1] = self.left_cnt[i]
        if left_arr:
            self.left = WaveletTree(left_arr, lo, mid)
        if right_arr:
            self.right = WaveletTree(right_arr, mid+1, hi)

    def kth(self, l, r, k):
        if self.lo == self.hi:
            return self.lo
        cnt_l = self.left_cnt[l]
        cnt_r = self.left_cnt[r]
        left_cnt = cnt_r - cnt_l
        if k <= left_cnt:
            return self.left.kth(cnt_l, cnt_r, k)
        else:
            return self.right.kth(l - cnt_l, r - cnt_r, k - left_cnt)

    def count_leq(self, l, r, v):
        if v < self.lo:
            return 0
        if self.hi <= v:
            return r - l
        if self.lo == self.hi:
            return r - l if self.lo <= v else 0
        cnt_l = self.left_cnt[l]
        cnt_r = self.left_cnt[r]
        result = 0
        if self.left:
            result += self.left.count_leq(cnt_l, cnt_r, v)
        if self.right:
            result += self.right.count_leq(l - cnt_l, r - cnt_r, v)
        return result

def main():
    N, Q = map(int, input().split())
    A = list(map(int, input().split()))
    sorted_vals = sorted(set(A))
    compress = {v: i for i, v in enumerate(sorted_vals)}
    A_comp = [compress[v] for v in A]
    wt = WaveletTree(A_comp, 0, len(sorted_vals)-1)
    results = []
    for _ in range(Q):
        q = input().split()
        if q[0] == 'kth':
            l, r, k = int(q[1])-1, int(q[2]), int(q[3])
            idx = wt.kth(l, r, k)
            results.append(sorted_vals[idx])
        else:
            l, r, v = int(q[1])-1, int(q[2]), int(q[3])
            vi = bisect_right(sorted_vals, v) - 1
            if vi < 0:
                results.append(0)
            else:
                results.append(wt.count_leq(l, r, vi))
    print('\n'.join(map(str, results)))

main()

Step-by-Step 解説

1構造
値域 [lo, hi] を mid で分割。各要素を左右に振り分け、prefix sum を持つ。
2k 番目
区間内の左行き数を見て、k ≤ なら左、それ以外は右で k-left。
3以下カウント
v が hi 以下なら全体、lo 未満なら 0、それ以外は左右再帰。
4座標圧縮
10^9 を log N の深さに抑える。

よくあるミス

ミス原因正しい書き方
インデックスオフセット0/1-indexed 混在入力時に変換
left_cnt の範囲外N+1 が必要[0]*(n+1)
座標復元忘れsorted_vals[idx]

次のステップ

  • 区間 distinct 要素数の数え上げ(オフライン)

自己評価

自分の回答

気づき・メモ