Day 111-Q2 — Scapegoat Tree(スケープゴート木・重み平衡BSTの償却再構築)

2026-08-03 赤色 Master / Phase 8+ ★★★★★★★★★ 部分木の完全再構築による償却 O(log N) の順序統計木

問題

空の集合に対して $Q$ 個のクエリを順に処理せよ。0 x(値 $x$ を挿入、相異なることが保証される)と 1 k($k$ 番目に小さい値を出力)の2種類。

入力形式

Q
query_1
...
query_Q

制約

$1 \le Q \le 2\times10^5$
$0 \le x \le 10^9$
1 k実行時、要素数は$k$以上

入出力例

入力例1

8
0 50
0 20
0 80
0 10
0 30
1 3
0 5
1 1

出力例1

30
5

$\{50,20,80,10,30\}$昇順$10,20,30,50,80$の3番目は$30$。$5$挿入後の1番目は$5$。

概念図: 部分木の flatten → 平衡再構築

偏った部分木(スケープゴート) 10 20 30 40 深さがO(log N)を超えた→再構築 flatten 完全平衡木として再構築 30 10 40 20 深さ O(log N) に復帰

ヒント(段階的開示)

ヒント1: 方向性
順序統計を扱う二分探索木は、挿入順序次第で深さが$O(N)$まで偏ることがある。回転で平衡を保つ方法もあるが、もっと単純に「偏った部分木ごと配列に展開して完全に再構築する」方式が使える。
ヒント2: アプローチ
通常のBST挿入後、深さが $\log_{1/\alpha} N$($\alpha\approx0.7$)を超えたら、挿入経路上で「子の部分木サイズが親の$\alpha$倍を超える」最初の祖先(スケープゴート)を見つけ、その部分木を中順展開→平衡再構築で置き換える。再構築コストは「そこに$\Theta(s)$回挿入が積もった後」にしか発生しないため償却$O(\log N)$。
ヒント3: 誘導(コード骨格)
def insert(root, val):
    # BST挿入しながら経路pathを記録
    # path を reversed() して size を下から上へ更新(重要)
    # depth > log(len(path))/log(1/ALPHA) ならスケープゴートを探索し
    # flatten(部分木) → build(平衡木) で置き換える

模範解答 (Python)

import sys, math

ALPHA = 0.7

class Node:
    __slots__ = ("val", "left", "right", "size")
    def __init__(self, val):
        self.val = val
        self.left = None
        self.right = None
        self.size = 1

def size(n):
    return n.size if n else 0

def update(n):
    n.size = 1 + size(n.left) + size(n.right)

def flatten_iter(root):
    out = []
    stack = []
    node = root
    while stack or node:
        while node:
            stack.append(node)
            node = node.left
        node = stack.pop()
        out.append(node.val)
        node = node.right
    return out

def build_iter(vals):
    n = len(vals)
    if n == 0:
        return None
    root = None
    stack = [(0, n - 1, None, None)]
    while stack:
        lo, hi, parent, side = stack.pop()
        mid = (lo + hi) // 2
        node = Node(vals[mid])
        if parent is None:
            root = node
        elif side == 'L':
            parent.left = node
        else:
            parent.right = node
        if lo <= mid - 1:
            stack.append((lo, mid - 1, node, 'L'))
        if mid + 1 <= hi:
            stack.append((mid + 1, hi, node, 'R'))
    stack2 = [root]
    order = []
    while stack2:
        nd = stack2.pop()
        if nd is None:
            continue
        order.append(nd)
        stack2.append(nd.left)
        stack2.append(nd.right)
    for nd in reversed(order):
        update(nd)
    return root

class ScapegoatTree:
    def __init__(self):
        self.root = None

    def insert(self, val):
        path = []
        node = self.root
        parent, side = None, None
        while node:
            path.append(node)
            if val < node.val:
                parent, side, node = node, 'L', node.left
            else:
                parent, side, node = node, 'R', node.right
        newnode = Node(val)
        if parent is None:
            self.root = newnode
        elif side == 'L':
            parent.left = newnode
        else:
            parent.right = newnode
        path.append(newnode)
        for n in reversed(path):
            update(n)

        depth = len(path) - 1
        max_depth = math.log(len(path)) / math.log(1 / ALPHA) if len(path) > 1 else 0
        if depth > max_depth:
            for i in range(len(path) - 2, -1, -1):
                n = path[i]
                if size(n.left) > ALPHA * size(n) or size(n.right) > ALPHA * size(n):
                    vals = flatten_iter(n)
                    newsub = build_iter(vals)
                    if i == 0:
                        self.root = newsub
                    else:
                        par = path[i - 1]
                        if par.left is n:
                            par.left = newsub
                        else:
                            par.right = newsub
                    break

    def kth(self, k):
        node = self.root
        while node:
            ls = size(node.left)
            if k == ls + 1:
                return node.val
            elif k <= ls:
                node = node.left
            else:
                k -= ls + 1
                node = node.right
        return None

def solve():
    data = sys.stdin.read().split()
    idx = 0
    q = int(data[idx]); idx += 1
    t = ScapegoatTree()
    out = []
    for _ in range(q):
        typ = data[idx]; idx += 1
        if typ == "0":
            x = int(data[idx]); idx += 1
            t.insert(x)
        else:
            k = int(data[idx]); idx += 1
            out.append(str(t.kth(k)))
    print("\n".join(out))

solve()
計算量: insert・kthともに償却 $O(\log N)$。$Q=2\times10^5$のランダム混在クエリでsorted(list)参照実装と一致を確認済み。$10^5$件挿入を約1秒で処理できることも実測済み(flatten/buildを非再帰化しているため深い偏りでも安全)。

Step-by-Step 解説

1BST挿入しながら経路を記録
根から挿入位置までの経路をpathに保存する。
2sizeを下から上へ更新
pathを新ノードから根へ辿って更新する。逆順にすると子のsizeが古いまま参照されてしまう。
3深さ超過でスケープゴートを探索
「子の部分木サイズが自分の$\alpha$倍を超える」最初の祖先を見つける。
4部分木を展開して再構築
中順展開したソート済み配列から完全平衡木を作り直す。
5kthを部分木サイズで二分探索
左部分木サイズ+1と$k$を比較しながら降りる。

よくあるミス

ミス原因正しい書き方
sizeを根から新ノードへ(上から下へ)更新子のsizeがまだ古い値を参照してしまうreversed()して末端から根へ更新する
スケープゴート判定の不等号を誤るバランス条件の意味を取り違える「子が親の$\alpha$倍より大きい」を検出条件にする
flatten/buildを再帰で書きRecursionError部分木サイズが最悪$O(N)$になりうる明示的スタックで非再帰にする
深さと「path要素数」を混同する深さ(辺数)とノード数(深さ+1)を取り違える両者を明確に区別して閾値式に当てはめる

次のステップ

  • 発展: 遅延削除フラグ+全体再構築による削除操作を追加する
  • 発展: $\alpha$を変えたときの挿入コストと実際の深さのトレードオフを実験する
  • 発展: Treap(Day024・Day061)との設計思想の違いを比較する

自己評価

自分の回答

気づき・メモ