問題
長さ 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 を持つ。
値域 [lo, hi] を mid で分割。各要素を左右に振り分け、prefix sum を持つ。
2k 番目
区間内の左行き数を見て、k ≤ なら左、それ以外は右で k-left。
区間内の左行き数を見て、k ≤ なら左、それ以外は右で k-left。
3以下カウント
v が hi 以下なら全体、lo 未満なら 0、それ以外は左右再帰。
v が hi 以下なら全体、lo 未満なら 0、それ以外は左右再帰。
4座標圧縮
10^9 を log N の深さに抑える。
10^9 を log N の深さに抑える。
よくあるミス
| ミス | 原因 | 正しい書き方 |
|---|---|---|
| インデックスオフセット | 0/1-indexed 混在 | 入力時に変換 |
| left_cnt の範囲外 | N+1 が必要 | [0]*(n+1) |
| 座標復元忘れ | sorted_vals[idx] |
次のステップ
- 区間 distinct 要素数の数え上げ(オフライン)