Day 058-Q4 — 永続セグメント木 + 区間k番目クエリ

2026-06-11 赤色 Master / Phase 8+ ★★★★★★★★★ Persistent Segment Tree / Merge Sort Tree / 座標圧縮 / 区間k番目

問題

長さ $N$ の整数列 $A_1, A_2, \ldots, A_N$ が与えられる。以下の $Q$ クエリを処理せよ。

  • クエリ 1 l r k: $A_l, \ldots, A_r$ の中で $k$ 番目に小さい値を出力せよ(1-indexed)
  • クエリ 2 l r x: $A_l, \ldots, A_r$ の中で $x$ 以下の要素の個数を出力せよ

制約

パラメータ範囲
$N$$1 \le N \le 2 \times 10^5$
$Q$$1 \le Q \le 2 \times 10^5$
$A_i$$1 \le A_i \le 10^9$
クエリ1 $k$$1 \le k \le r - l + 1$
クエリ2 $x$$1 \le x \le 10^9$

入出力例

入力例 1

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

出力例 1

3
4
4

クエリ1: $A[2..6]=[3,1,4,1,5]$ → ソート $[1,1,3,4,5]$ → 3番目 = 3。
クエリ2: $A[1..5]=[5,3,1,4,1]$ → 4以下の個数 = 4。
クエリ3: $A[1..7]=[5,3,1,4,1,5,9]$ → ソート → 4番目 = 4。

概念図: 永続セグメント木の差分クエリ

Persistent Segment Tree — バージョン差分で区間頻度分布を取得 roots[r]: A[1..r] の頻度 SegTree root_r cnt cnt roots[l-1]: A[1..l-1] の頻度 SegTree root_{l-1} cnt cnt 差分 = A[l..r] の頻度分布 root_r - root_{l-1} 差分木上で k番目を二分探索(SegTree Walk) left_cnt = nodes[v.left].cnt - nodes[u.left].cnt if k ≤ left_cnt: 左部分木に降りる else: 右部分木に降りる (k -= left_cnt) 座標圧縮: A_i ≤ 10^9 → [1, M] (M ≤ N) SegTree サイズを O(N) に削減。更新ごとに O(log N) ノードを新規生成。

ヒント(段階的開示)

ヒント1: 方向性
区間 $k$ 番目クエリには「永続セグメント木(Persistent Segment Tree)」を使うのが定石。 座標圧縮後、$i$ 番目の接頭辞に対応する永続 SegTree を構築し、 差分で区間クエリを $O(\log N)$ で処理する。
ヒント2: アプローチ
  1. $A$ の値を座標圧縮(rank 変換)
  2. 永続 SegTree の $i$ 番目バージョン = $A[1..i]$ の値の頻度を記録
  3. クエリ $[l,r]$ = バージョン $r$ − バージョン $l-1$ の差分 SegTree
  4. 差分木上で k 番目を二分探索(SegTree Walk): $O(\log N)$
ヒント3: コード骨格
def update(prev, l, r, pos):
    """pos に +1 した新しいノードを返す"""
    cur = new_node()
    nodes[cur].cnt = nodes[prev].cnt + 1
    nodes[cur].left = nodes[prev].left
    nodes[cur].right = nodes[prev].right
    if l == r:
        return cur
    mid = (l + r) >> 1
    if pos <= mid:
        nodes[cur].left = update(nodes[prev].left, l, mid, pos)
    else:
        nodes[cur].right = update(nodes[prev].right, mid+1, r, pos)
    return cur

def query_kth(u, v, l, r, k):
    if l == r:
        return l
    mid = (l + r) >> 1
    left_cnt = nodes[nodes[v].left].cnt - nodes[nodes[u].left].cnt
    if k <= left_cnt:
        return query_kth(nodes[u].left, nodes[v].left, l, mid, k)
    return query_kth(nodes[u].right, nodes[v].right, mid+1, r, k - left_cnt)

模範解答 (Python)

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

class Node:
    __slots__ = ['left', 'right', 'cnt']
    def __init__(self):
        self.left = self.right = 0
        self.cnt = 0

MAXN = 200005 * 40
nodes = [Node() for _ in range(MAXN)]
node_cnt = [0]

def new_node():
    node_cnt[0] += 1
    return node_cnt[0]

def update(prev, l, r, pos):
    cur = new_node()
    nodes[cur].cnt = nodes[prev].cnt + 1
    nodes[cur].left = nodes[prev].left
    nodes[cur].right = nodes[prev].right
    if l == r:
        return cur
    mid = (l + r) >> 1
    if pos <= mid:
        nodes[cur].left = update(nodes[prev].left, l, mid, pos)
    else:
        nodes[cur].right = update(nodes[prev].right, mid+1, r, pos)
    return cur

def query_kth(u, v, l, r, k):
    if l == r:
        return l
    mid = (l + r) >> 1
    lc = nodes[nodes[v].left].cnt - nodes[nodes[u].left].cnt
    if k <= lc:
        return query_kth(nodes[u].left, nodes[v].left, l, mid, k)
    return query_kth(nodes[u].right, nodes[v].right, mid+1, r, k - lc)

def query_leq(u, v, l, r, x):
    if r <= x:
        return nodes[v].cnt - nodes[u].cnt
    if l > x:
        return 0
    mid = (l + r) >> 1
    return (query_leq(nodes[u].left, nodes[v].left, l, mid, x) +
            query_leq(nodes[u].right, nodes[v].right, mid+1, r, x))

def main():
    N, Q = map(int, input().split())
    A = list(map(int, input().split()))
    sorted_vals = sorted(set(A))
    M = len(sorted_vals)
    compress = {v: i+1 for i, v in enumerate(sorted_vals)}
    A_comp = [compress[a] for a in A]

    roots = [0] * (N + 1)
    for i in range(N):
        roots[i+1] = update(roots[i], 1, M, A_comp[i])

    results = []
    for _ in range(Q):
        line = input().split()
        t = line[0]
        if t == '1':
            l, r, k = int(line[1])-1, int(line[2]), int(line[3])
            rank = query_kth(roots[l], roots[r], 1, M, k)
            results.append(sorted_vals[rank-1])
        else:
            l, r, x = int(line[1])-1, int(line[2]), int(line[3])
            xr = bisect_right(sorted_vals, x)
            if xr == 0:
                results.append(0)
            else:
                cnt = query_leq(roots[l], roots[r], 1, M, xr)
                results.append(cnt)

    print('\n'.join(map(str, results)))

main()

Step-by-Step 解説

Step 1: 座標圧縮

$A_i \le 10^9$ を $[1, M]$($M \le N$)に圧縮。SegTree のサイズを $O(N)$ に削減する。

Step 2: 永続 SegTree の構築

ルート配列 roots[0..N]roots[i] = $A[1..i]$ の値頻度を管理する SegTree のルートノード。各更新で変化したノードのみ新規生成(パス上 $O(\log N)$ ノード)。

Step 3: 差分クエリ

$[l, r]$ の頻度分布 = roots[r] - roots[l-1]。差分木上で二分探索 → k 番目 $O(\log N)$。

Step 4: x 以下のカウント

座標圧縮後の x の rank を bisect_right で求め、左部分木のカウントを再帰的に合計する。

よくあるミス

ミス原因正しい書き方
ノード数不足永続化で最大 O(N log N) ノード必要MAXN = N * 40 程度確保
座標圧縮の境界x が配列に存在しないbisect_right で rank を取得
1-indexed 混乱l, r の変換ミスクエリは l-1, r で区間 $[l-1, r)$

次のステップ

発展問題: 点更新($A_i$ の値を変更)しながら同じクエリを処理せよ(動的 Merge Sort Tree)。

自己評価

解いた後に記入してください。