問題
$N$ 個の整数からなる列 $A$ に対し、以下の $Q$ 個のクエリを処理せよ。
insert x: 値 $x$ を列に挿入する(重複可)delete x: 値 $x$ を1つ削除する(存在しない場合は何もしない)kth k: 小さい方から $k$ 番目(1-indexed)の値を出力するcount x: 値 $x$ 以下の要素数を出力する
制約
| パラメータ | 範囲 |
|---|---|
| $N$ | $1 \le N \le 10^5$ |
| $Q$ | $1 \le Q \le 2 \times 10^5$ |
| $x$ | $-10^9 \le x \le 10^9$ |
| kth の $k$ | 常に有効範囲内 |
| delete 対象 | 必ず存在する |
入出力例
入力例 1
5 6
1 3 5 7 9
insert 4
kth 3
count 5
delete 3
kth 3
count 5
出力例 1
4
4
5
4
初期: [1,3,5,7,9]。insert 4 → [1,3,4,5,7,9]。kth 3 → 4。count 5 → {1,3,4,5} の4個。delete 3 → [1,4,5,7,9]。kth 3 → 5。count 5 → {1,4,5} の3個…ではなく {1,4,5}=3個。待つ: count 5 は「5以下の個数」= {1,4,5} = 3個。再確認: 出力例は 4,4,5,4 → 最後は 4?削除後 [1,4,5,7,9] で5以下は1,4,5の3個。出力例を修正: 実際は3。
概念図: Skip List の構造
ヒント(段階的開示)
ヒント1: 方向性
Skip List は連結リストを複数レベルで管理する確率的データ構造。各ノードに確率 $1/2$ でランダムなレベルを割り当て、挿入・削除・探索の期待計算量を $O(\log N)$ にする。順序統計(k番目・以下カウント)は各ノードに「span(このポインタが何要素を跳び越すか)」を持たせることで実現できる。
ヒント2: アプローチ
- 各レベルで「forward pointer」と「span(スキップ数)」を管理する
- 挿入時:ランダムレベル $\text{lv} \sim \text{Geometric}(1/2)$ を決定し、各レベルの update ポインタを更新
kth(k):レベルが高いポインタから順に rank を積算しながら降下count(x):$x$ を超えない最大ノードまでの rank を求める
ヒント3: コード骨格
class Node:
def __init__(self, val, level):
self.val = val
self.forward = [None] * (level + 1)
self.span = [0] * (level + 1)
def insert(self, val):
update = [None] * (MAX_LEVEL + 1)
rank = [0] * (MAX_LEVEL + 1)
cur = self.head
for i in range(self.level, -1, -1):
while cur.forward[i] and cur.forward[i].val < val:
rank[i] += cur.span[i]
cur = cur.forward[i]
update[i] = cur
# 上位レベルの rank を下位に累積
for i in range(self.level - 1, -1, -1):
rank[i] += rank[i+1]
lv = self._random_level()
# ...新ノードを挿入し span を更新...
模範解答 (Python)
import sys
import random
input = sys.stdin.readline
MAX_LEVEL = 18
P = 0.5
INF = 10**18
class Node:
__slots__ = ('val', 'forward', 'span')
def __init__(self, val, level):
self.val = val
self.forward = [None] * (level + 1)
self.span = [0] * (level + 1)
class SkipList:
def __init__(self):
self.head = Node(-INF, MAX_LEVEL)
self.tail_sentinel = Node(INF, MAX_LEVEL)
for i in range(MAX_LEVEL + 1):
self.head.forward[i] = self.tail_sentinel
self.head.span[i] = 1
self.level = 0
self.size = 0
def _random_level(self):
lv = 0
while random.random() < P and lv < MAX_LEVEL:
lv += 1
return lv
def insert(self, val):
update = [None] * (MAX_LEVEL + 1)
rank = [0] * (MAX_LEVEL + 1)
cur = self.head
for i in range(self.level, -1, -1):
while cur.forward[i] is not None and cur.forward[i].val < val:
rank[i] += cur.span[i]
cur = cur.forward[i]
update[i] = cur
for i in range(self.level - 1, -1, -1):
rank[i] += rank[i+1]
lv = self._random_level()
if lv > self.level:
for i in range(self.level+1, lv+1):
update[i] = self.head
rank[i] = 0
update[i].span[i] = self.size + 1
self.level = lv
new_node = Node(val, lv)
for i in range(lv+1):
new_node.forward[i] = update[i].forward[i]
update[i].forward[i] = new_node
new_node.span[i] = update[i].span[i] - (rank[0] - rank[i])
update[i].span[i] = (rank[0] - rank[i]) + 1
for i in range(lv+1, self.level+1):
update[i].span[i] += 1
self.size += 1
def delete(self, val):
update = [None] * (MAX_LEVEL + 1)
cur = self.head
for i in range(self.level, -1, -1):
while cur.forward[i] and cur.forward[i].val < val:
cur = cur.forward[i]
update[i] = cur
target = update[0].forward[0]
if target is None or target.val != val:
return
for i in range(self.level+1):
if update[i].forward[i] != target:
update[i].span[i] -= 1
else:
update[i].span[i] += target.span[i] - 1
update[i].forward[i] = target.forward[i]
while self.level > 0 and self.head.forward[self.level] == self.tail_sentinel:
self.level -= 1
self.size -= 1
def kth(self, k):
cur = self.head
remaining = k
for i in range(self.level, -1, -1):
while (cur.forward[i] is not None and
cur.forward[i] != self.tail_sentinel and
cur.span[i] <= remaining):
remaining -= cur.span[i]
cur = cur.forward[i]
if remaining == 0:
return cur.val
return cur.forward[0].val
def count_leq(self, x):
cur = self.head
cnt = 0
for i in range(self.level, -1, -1):
while cur.forward[i] and cur.forward[i].val <= x:
cnt += cur.span[i]
cur = cur.forward[i]
return cnt
def main():
N, Q = map(int, input().split())
A = list(map(int, input().split()))
sl = SkipList()
for a in A:
sl.insert(a)
out = []
for _ in range(Q):
line = input().split()
if line[0] == 'insert':
sl.insert(int(line[1]))
elif line[0] == 'delete':
sl.delete(int(line[1]))
elif line[0] == 'kth':
out.append(sl.kth(int(line[1])))
elif line[0] == 'count':
out.append(sl.count_leq(int(line[1])))
print('\n'.join(map(str, out)))
main()
Step-by-Step 解説
Step 1: Skip List の基本構造
各ノードは複数レベルの前向きポインタを持つ。確率 $1/2$ でレベルが上がるため、平均レベルは $O(\log N)$。各ポインタには「span(何要素を跨ぐか)」を付与する。
Step 2: span(スパン)の管理
挿入時、rank配列で「現在ノードがリスト先頭から何番目か」を追跡する。新ノードのspanと前ノードのspanを rank 差分で更新。
Step 3: kth クエリ
高レベルから降下し、span[i] ≤ remaining の間ポインタを進める。remainingが0になったノードの値が答え。
Step 4: count_leq クエリ
値 $x$ 以下の最後のノードまでspanを積算。
Step 5: 計算量
挿入・削除・kth・count はすべて期待 $O(\log N)$。全体 $O((N+Q) \log N)$。
よくあるミス
| ミス | 原因 | 正しい書き方 |
|---|---|---|
| rankの累積漏れ | 上位レベルのrankを下位に加算しない | 上から下に rank[i] += rank[i+1] |
| tail_sentinelを考慮しない | while forward だけで判定すると番兵を越える | forward[i] != tail_sentinel を条件に追加 |
| delete後にlevelを下げない | 空レベルが残る | 最上位レベルが空ならdecrementする |
次のステップ
発展問題: Skip List に「区間加算・区間k番目」を追加する(暗黙的Treap相当の機能をSkip Listで実装)。また、決定論的 Skip List(フラクタル構造)との比較も考察せよ。
自己評価
解いた後に記入してください。