Day 105-Q4 — Wavelet Tree(区間×値域カウント)

2026-07-28 赤色 Master / Phase 8+ ★★★★★★★★☆ 値方向の二分木・prefix countによる区間伝播

問題

長さ$N$の数列$A=(A_1,\dots,A_N)$が与えられる。$Q$個のクエリ$(l_i,r_i,x_i,y_i)$が与えられ、区間$[l_i,r_i]$(1-indexed、両端含む)に含まれる要素のうち、値が$x_i$以上$y_i$以下であるものの個数を答えよ。

この「区間×値域」の2次元カウントは Merge Sort Tree でも$O((\log N)^2)$程度で解けるが、値の範囲でノードを分割していくWavelet Treeを使うと$O(\log(\max A))$で処理でき、「区間内でk番目に小さい値」のような順序統計クエリにも同じ木で答えられる。各ノードは「値が左半分か右半分か」を表す0/1列の累積和(prefix count)を持ち、それでインデックス区間を子ノードに伝播する。

入力形式

N
A_1 A_2 ... A_N
Q
l_1 r_1 x_1 y_1
...
l_Q r_Q x_Q y_Q

制約

$1 \le N,Q \le 2\times10^5$
$-10^9 \le A_i \le 10^9$
$1 \le l_i \le r_i \le N$
$-10^9 \le x_i \le y_i \le 10^9$

入出力例

入力例1

8
5 2 9 2 5 3 8 1
3
1 8 2 5
2 6 1 3
4 8 5 9

出力例1

5
3
2

区間$[1,8]=\{5,2,9,2,5,3,8,1\}$のうち$[2,5]$は$5,2,2,5,3$の5個。区間$[2,6]=\{2,9,2,5,3\}$のうち$[1,3]$は$2,2,3$の3個。区間$[4,8]=\{2,5,3,8,1\}$のうち$[5,9]$は$5,8$の2個。

概念図: 値方向の分割とprefix countによる伝播

ノードの値域[lo,hi]をmidで左右に分割。位置iは prefix[i] で子の位置に写像 [lo,hi] ルート: ランク列全体 [lo,mid] 左: v ≤ mid [mid+1,hi] 右: v > mid [l',r') = [prefix[l], prefix[r]) [l',r') = [l-prefix[l], r-prefix[r])

ヒント(段階的開示)

ヒント1: 方向性
愚直に数えると$O(N)$かかりQ個のクエリでは間に合わない。セグメント木を「インデックス方向」に分割する発想(Merge Sort Tree)とは逆に、「値方向」に分割する木を考えよう。値の中央値`mid`を境に左半分・右半分に元の順序を保ったまま分岐していくイメージ。
ヒント2: アプローチ
座標圧縮したランク空間$[0,m)$上に木を再帰構築する。各ノードで要素列を先頭から見て「ランクがmid以下なら左、そうでなければ右」に振り分け、「先頭からi個のうち何個が左に行ったか」の累積和(prefix count)を持たせる。区間クエリ$[l,r)$はこのprefixで子ノードの対応区間に変換しながら木を降りる。
ヒント3: 誘導(コード骨格)
def build(vals, lo, hi):
    if lo == hi:
        return Node(lo, hi, None, None, None)
    mid = (lo + hi) // 2
    prefix = [0] * (len(vals) + 1)
    left_vals, right_vals = [], []
    for i, v in enumerate(vals):
        goes_left = (v <= mid)
        prefix[i+1] = prefix[i] + (1 if goes_left else 0)
        (left_vals if goes_left else right_vals).append(v)
    return Node(lo, hi, prefix, build(left_vals,lo,mid), build(right_vals,mid+1,hi))

def query(node, l, r, x, y):
    if l >= r or node.hi < x or node.lo > y: return 0
    if x <= node.lo and node.hi <= y: return r - l
    ll, rr = node.prefix[l], node.prefix[r]
    return query(node.left, ll, rr, x, y) + query(node.right, l-ll, r-rr, x, y)

模範解答 (Python)

import sys
import bisect


def solve():
    data = sys.stdin.buffer.read().split()
    idx = 0
    n = int(data[idx]); idx += 1
    a = list(map(int, data[idx:idx + n])); idx += n
    q = int(data[idx]); idx += 1
    queries = []
    for _ in range(q):
        l = int(data[idx]) - 1; idx += 1
        r = int(data[idx]); idx += 1
        x = int(data[idx]); idx += 1
        y = int(data[idx]); idx += 1
        queries.append((l, r, x, y))

    sys.setrecursionlimit(500000)

    sorted_vals = sorted(set(a))
    comp = {v: i for i, v in enumerate(sorted_vals)}
    ranks = [comp[v] for v in a]
    m = len(sorted_vals)

    class Node:
        __slots__ = ('lo', 'hi', 'prefix', 'left', 'right')

        def __init__(self, lo, hi, prefix, left, right):
            self.lo = lo
            self.hi = hi
            self.prefix = prefix
            self.left = left
            self.right = right

    def build(vals, lo, hi):
        if lo == hi:
            return Node(lo, hi, None, None, None)
        mid = (lo + hi) // 2
        prefix = [0] * (len(vals) + 1)
        left_vals = []
        right_vals = []
        for i, v in enumerate(vals):
            goes_left = v <= mid
            prefix[i + 1] = prefix[i] + (1 if goes_left else 0)
            if goes_left:
                left_vals.append(v)
            else:
                right_vals.append(v)
        return Node(lo, hi, prefix,
                     build(left_vals, lo, mid),
                     build(right_vals, mid + 1, hi))

    root = build(ranks, 0, m - 1) if m > 0 else None

    def query(node, l, r, x, y):
        if node is None or l >= r or node.hi < x or node.lo > y:
            return 0
        if x <= node.lo and node.hi <= y:
            return r - l
        left_l = node.prefix[l]
        left_r = node.prefix[r]
        cnt = query(node.left, left_l, left_r, x, y)
        right_l = l - left_l
        right_r = r - left_r
        cnt += query(node.right, right_l, right_r, x, y)
        return cnt

    out = []
    for l, r, x, y in queries:
        lo_rank = bisect.bisect_left(sorted_vals, x)
        hi_rank = bisect.bisect_right(sorted_vals, y) - 1
        if lo_rank > hi_rank:
            out.append('0')
            continue
        out.append(str(query(root, l, r, lo_rank, hi_rank)))

    print('\n'.join(out))


solve()
計算量: 構築$O(N\log m)$、各クエリ$O(\log m)$($m$=相異なる値の個数)。ランダム300ケースで区間内の値を全走査する愚直解との一致を確認済み。

Step-by-Step 解説

1座標圧縮
$A_i\le10^9$なので値をランク付けし、圧縮後の空間$[0,m)$で木を構築。クエリの値範囲もbisectでランク範囲に変換する。
2木の構築(値方向の分割)
各ノードは値のランク範囲$[lo,hi]$と要素列(元の相対順序を保持)を持つ。中央値`mid`で左右に振り分け累積和`prefix`を保存する。
3クエリのインデックス変換
`prefix[l]`は「位置lより前で左に振り分けられた要素数」。左の子の対応区間は$[prefix[l],prefix[r])$、右の子は$[l-prefix[l], r-prefix[r])$。
4枝刈りによる早期終了
ノードの値域がクエリの値範囲に完全に含まれれば$r-l$を即返す。全く重ならなければ$0$を返す。

よくあるミス

ミス原因正しい書き方
クエリの値$x,y$をそのままランクとして扱う木はランク空間で構築されているため対応が取れない`bisect_left`/`bisect_right`でランク範囲に変換してから問い合わせる
$A$に存在しない値だけの範囲の処理漏れ変換後に`lo_rank > hi_rank`となり不正な範囲でクエリする変換後にその条件をチェックし即座に$0$を返す
インデックス方向に分割してMerge Sort Treeと混同Wavelet Treeは値方向、Merge Sort Treeはインデックス方向という本質的な違いを見落とす分割の中央値は常に値(ランク)の中央であることを意識する
再帰の深さ制限に引っかかるランク空間の大きさに応じて再帰が深くなる`sys.setrecursionlimit`を大きく設定する

次のステップ

  • 発展: 同じ木構造で「区間内でk番目に小さい値」も$O(\log(\max A))$で求められることを確認する
  • 発展: succinctなビットベクトル(rank/select $O(1)$)を使い、メモリを$O(N\log(\max A))$ビットまで削減する本格実装に挑戦する
  • 次回予告: Karp's Minimum Mean Cycle(最小平均閉路検出・$O(VE)$のDP)

自己評価

自分の回答

気づき・メモ