Day 115-Q1 — スプレー木(Splay Tree・自己調整二分探索木による順序統計クエリ)

2026-08-07 赤色 Master / Phase 8+ ★★★★★★★★★ 自己調整平衡二分探索木・償却計算量解析

問題

整数の多重集合(同じ値を複数個含められる集合)に対して、$Q$ 個のクエリを処理せよ。クエリは次の4種類である。

  • 1 x — 値 $x$ を1個追加する
  • 2 x — 値 $x$ を1個削除する(その時点で $x$ が少なくとも1個存在することが保証される)
  • 3 x — 現在の多重集合の中で $x$ より小さい値の個数を数え、その個数 $+1$(すなわち $x$ を追加したときの順位)を出力する
  • 4 k — 現在の多重集合の中で $k$ 番目に小さい値($1$-indexed。同じ値は多重度に応じて別々に数える)を出力する

入力形式

Q
query_1
query_2
...
query_Q

制約

$1 \le Q \le 2\times10^5$
$1 \le x \le 10^9$
クエリ4 kの時点で $k$ は多重集合のサイズ以下

入出力例

入力例1

7
1 5
1 3
1 8
3 5
4 1
2 3
4 1

出力例1

2
3
5

1 5,1 3,1 8で多重集合は$\{3,5,8\}$。3 5は$5$未満の個数$1$($3$のみ)$+1=2$。4 1は最小値$3$。2 3で$3$を削除し$\{5,8\}$。4 1は最小値$5$。

概念図: zig-zig と zig-zag の回転パターン

アクセスしたノードxを根まで持ち上げる2種類の二重回転 zig-zig(同じ向き) g p x rotate(p) → rotate(x) の順に2回回転 x p g zig-zag(逆向き) g p x rotate(x) → rotate(x) 同じ関数を2回連続 zig-zigは「親を先に上げる」、zig-zagは「xを2回連続で上げる」— この違いが 最悪ケースでも償却 $O(\log N)$ を保証するポテンシャル解析の鍵になる

ヒント(段階的開示)

ヒント1: 方向性
配列をソート済みに保って挿入位置を探す方法では、挿入・削除のたびに要素シフトで $O(N)$ かかり間に合わない。平衡二分探索木を使えば挿入・削除・順位・$k$番目探索をすべて $O(\log N)$ 程度で行えるが、AVL木や赤黒木は回転条件の管理が複雑になりがち。もっと単純な「アクセスした要素を根まで持ち上げる」だけの木を考えてみる。
ヒント2: アプローチ
スプレー木は明示的なバランス条件を持たない代わりに、ノードにアクセスするたびに「スプレー操作」でそのノードを回転を繰り返して根まで引き上げる。zig(親が根)・zig-zig(親と祖父が同じ向きの子)・zig-zag(逆向きの子)の3パターンの回転を使い分けることで、償却計算量が $O(\log N)$ になることが保証される(ポテンシャル関数によるならし解析)。挿入は「木を降りて挿入位置の直前・直後に来るノードを見つけてスプレーで根にしてから分割挿入する」、削除は「削除対象をスプレーで根にしてから左右の部分木をマージする」という形で実装できる。
ヒント3: 誘導(コード骨格)
def rotate(x): # xを親pの位置に持ち上げる単一回転
    ...
def splay(x): # par[x]!=0の間、zig-zig/zig-zag/zigを判定しrotateを繰り返す
    while par[x]:
        p = par[x]; g = par[p]
        if g:
            if (left[g]==p) == (left[p]==x):
                rotate(p)   # zig-zig: 親を先に上げる
            else:
                rotate(x)   # zig-zag: xを先に上げる
        rotate(x)
    root = x

模範解答 (Python)

import sys

def main():
    data = sys.stdin.buffer.read().split()
    idx = 0
    Q = int(data[idx]); idx += 1

    MAXN = Q + 5
    key = [0] * MAXN
    cnt = [0] * MAXN
    sz = [0] * MAXN
    left = [0] * MAXN
    right = [0] * MAXN
    par = [0] * MAXN
    NIL = 0
    n_nodes = 0
    root = NIL

    def new_node(x, p):
        nonlocal n_nodes
        n_nodes += 1
        key[n_nodes] = x
        cnt[n_nodes] = 1
        sz[n_nodes] = 1
        left[n_nodes] = NIL
        right[n_nodes] = NIL
        par[n_nodes] = p
        return n_nodes

    def pull(t):
        sz[t] = cnt[t] + sz[left[t]] + sz[right[t]]

    def rotate(x):
        p = par[x]
        g = par[p]
        if left[p] == x:
            left[p] = right[x]
            if right[x]:
                par[right[x]] = p
            right[x] = p
        else:
            right[p] = left[x]
            if left[x]:
                par[left[x]] = p
            left[x] = p
        par[p] = x
        par[x] = g
        if g:
            if left[g] == p:
                left[g] = x
            else:
                right[g] = x
        pull(p)
        pull(x)

    def splay(x):
        nonlocal root
        while par[x]:
            p = par[x]
            g = par[p]
            if g:
                if (left[g] == p) == (left[p] == x):
                    rotate(p)
                else:
                    rotate(x)
            rotate(x)
        root = x

    def find(x):
        nonlocal root
        t = root
        last = NIL
        while t:
            last = t
            if key[t] == x:
                break
            elif x < key[t]:
                t = left[t]
            else:
                t = right[t]
        if last:
            splay(last)
        return last != NIL and key[last] == x

    def insert(x):
        nonlocal root
        if root == NIL:
            root = new_node(x, 0)
            return
        found = find(x)
        if found:
            cnt[root] += 1
            pull(root)
            return
        newr = new_node(x, 0)
        if x < key[root]:
            left[newr] = left[root]
            if left[root]:
                par[left[root]] = newr
            right[newr] = root
            left[root] = 0
            par[root] = newr
        else:
            right[newr] = right[root]
            if right[root]:
                par[right[root]] = newr
            left[newr] = root
            right[root] = 0
            par[root] = newr
        pull(root)
        root = newr
        pull(root)

    def delete(x):
        nonlocal root
        find(x)
        if cnt[root] > 1:
            cnt[root] -= 1
            pull(root)
            return
        l = left[root]
        r = right[root]
        if l:
            par[l] = 0
        if r:
            par[r] = 0
        if l == 0:
            root = r
        elif r == 0:
            root = l
        else:
            t = l
            while right[t]:
                t = right[t]
            splay(t)
            right[t] = r
            if r:
                par[r] = t
            pull(t)
            root = t
        if root:
            par[root] = 0

    def kth(k):
        nonlocal root
        t = root
        while True:
            lsz = sz[left[t]]
            if k <= lsz:
                t = left[t]
            elif k <= lsz + cnt[t]:
                splay(t)
                return key[t]
            else:
                k -= lsz + cnt[t]
                t = right[t]

    def rank(x):
        t = root
        r = 0
        last = NIL
        while t:
            last = t
            if key[t] < x:
                r += sz[left[t]] + cnt[t]
                t = right[t]
            else:
                t = left[t]
        if last:
            splay(last)
        return r + 1

    out = []
    for _ in range(Q):
        c = data[idx]; idx += 1
        if c == b'1':
            x = int(data[idx]); idx += 1
            insert(x)
        elif c == b'2':
            x = int(data[idx]); idx += 1
            delete(x)
        elif c == b'3':
            x = int(data[idx]); idx += 1
            out.append(str(rank(x)))
        else:
            k = int(data[idx]); idx += 1
            out.append(str(kth(k)))

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

main()
計算量: 挿入・削除・順位・$k$番目探索いずれも償却 $O(\log N)$(スプレー操作のならし解析による)。全体で $O(Q \log Q)$。$Q=2\times10^5$のランダムクエリ列で、`sortedcontainers`ベースの参照実装と全出力が一致することを確認済み。

Step-by-Step 解説

1rotate — 単一回転の基本操作
$x$ をその親 $p$ の位置に持ち上げる。$x$ が $p$ の左の子なら右回転、右の子なら左回転。祖父 $g$ の子ポインタも $x$ に張り替える。pullは必ず「子→親」の順で呼ぶ。
2splay — zig / zig-zig / zig-zag の使い分け
祖父がなければ単純な回転(zig)を1回。同じ向きの親子関係(zig-zig)ならrotate(p)rotate(x)。逆向き(zig-zag)ならrotate(x)を2回。順序を間違えると償却計算量の保証が崩れる。
3find(x) による insert の実装
findは見つかった(または最後に訪れた)ノードをスプレーして根にする。挿入時は根とその大小関係で左右どちらかの子と根を差し替える2分岐挿入を行う。
4delete(x) — 根にしてから左右部分木を統合
削除対象を根にし(多重度2以上ならカウントを減らすだけ)、左部分木の最大値ノードをスプレーして新しい根にし、そこへ右部分木をぶら下げる。
5kth(k)rank(x) — サイズフィールドと償却計算量の維持
sz[left[t]]で順位探索。アクセスしたノードは必ずsplayし、連続アクセスでも $O(\log N)$ を維持する。

よくあるミス

ミス原因正しい書き方
挿入時にsplayを呼ばず、木が偏ったまま蓄積する「見つけたら終わり」で満足してしまう挿入・探索・削除いずれの操作末尾でも対象ノードをsplayして根にする
deleteで左部分木の最大値を切り離す際にparを更新し忘れる部分木を切り離す処理を省略left[root]/right[root]を切り離す際は必ずpar[l]=0, par[r]=0を設定してからsplayする
zig-zigとzig-zagの回転順序を取り違える片方のパターンだけ意識して実装zig-zigはrotate(p)rotate(x)、zig-zagはrotate(x)を2回連続
kth/rankクエリでsplayを呼ばず償却計算量が崩れる「値を返すだけだから根に上げなくていい」と誤解アクセスしたノードは必ずsplayし、連続操作でも $O(\log N)$ を維持する

次のステップ

  • 発展: 区間reverse等、implicit treapのような「列としてのスプレー木」操作を追加してみる
  • 発展: split/mergeベースのスプレー木実装に書き換え、今回のsplay-then-linkスタイルとの計算量・実装量を比較する
  • 発展: 前回の永続Treapと同様に、スプレー木を永続化してみる(copy-on-writeが必要になる理由を考える)

自己評価

自分の回答

気づき・メモ