問題
長さ $N$ の数列 $A$ がある。この数列に対して $Q$ 個のクエリを バージョン管理しながら 処理せよ。
各クエリはバージョン番号 $v$($0$ は初期状態)を指定し、そのバージョンに対して操作を行い 新しいバージョン を作る。クエリは次の4種類である。
1 v l r— バージョン $v$ の数列の区間 $[l, r]$($1$-indexed, 両端含む)を反転し、新バージョンとする2 v l r— バージョン $v$ の数列の区間 $[l, r]$ の総和を出力する(新バージョンは作らない)3 v x y— バージョン $v$ の数列の $x$ 番目と $y$ 番目の値を交換し、新バージョンとする4 v— バージョン $v$ をそのまま複製し、新バージョンとする(過去の任意バージョンへの「分岐」の確認用)
新しいバージョンは常に、それまでに作られたバージョンの最大番号 $+1$ として採番される。
入力形式
N Q
A_1 A_2 ... A_N
query_1
query_2
...
query_Q
制約
$1 \le N, Q \le 2\times10^5$
$1 \le A_i \le 10^9$
$0 \le v <$ その時点で存在するバージョン数
$1 \le l \le r \le N$, $1 \le x, y \le N$
入出力例
入力例1
5 6
1 2 3 4 5
1 0 2 4
2 1 1 5
3 0 1 5
2 2 1 5
4 1
2 3 2 4
出力例1
15
15
9
バージョン0は1 2 3 4 5。バージョン1は区間[2,4]反転で1 4 3 2 5(全区間和15)。バージョン2は元のバージョン0で1番目と5番目を交換した5 2 3 4 1(全区間和15)。バージョン3はバージョン1の複製1 4 3 2 5で、その[2,4]の和は $4+3+2=9$。
概念図: コピー・オン・ライトで過去バージョンを保持する
ヒント(段階的開示)
ヒント1: 方向性
「バージョンごとに配列をまるごとコピーする」と1回の操作が $O(N)$ になり間に合わない。過去のバージョンを保持したまま更新前後で 差分だけ新しいノードを作る データ構造、いわゆる永続データ構造が必要になる。列の反転・区間和という操作は暗黙的Treap(implicit treap)で実現できるので、それを永続化する。
ヒント2: アプローチ
暗黙的Treapのノードを更新するとき、既存のノードを書き換えるのではなく 新しいノードを作って左右の子だけ差し替える(コピー・オン・ライトの考え方)。優先度は固定のまま乱数で決めておき、分割・併合のたびに通過した経路上のノードだけを複製すれば、1回の操作あたり $O(\log N)$ 個の新ノードで済む。反転は遅延伝播フラグ(reverse flag)を使うが、フラグを子に伝播する際にも子を複製してから書き換える。
ヒント3: 誘導(コード骨格)
def split(t, k): # tを先頭k個とそれ以降に分割。通過ノードは新規複製する。
...
def merge(a, b): # a, bを併合。通過ノードは新規複製する。
...
def push_down(t): # 遅延反転フラグがあれば子を複製してから交換する
def reverse(root, l, r):
a, bc = split(root, l - 1)
b, c = split(bc, r - l + 1)
b.rev ^= True
return merge(merge(a, b), c) # 新バージョンとして保存
模範解答 (Python)
import sys, random
def main():
input_data = sys.stdin.buffer.read().split()
idx = 0
N, Q = int(input_data[idx]), int(input_data[idx+1]); idx += 2
A = [int(input_data[idx+i]) for i in range(N)]; idx += N
random.seed(12345)
val = [0]; sm = [0]; sz = [0]; pr = [0]; lc = [0]; rc = [0]; rev = [False]
NIL = 0
def new_node(v):
val.append(v); sm.append(v); sz.append(1); pr.append(random.random())
lc.append(NIL); rc.append(NIL); rev.append(False)
return len(val) - 1
def clone(t):
val.append(val[t]); sm.append(sm[t]); sz.append(sz[t]); pr.append(pr[t])
lc.append(lc[t]); rc.append(rc[t]); rev.append(rev[t])
return len(val) - 1
def pull(t):
sm[t] = val[t] + sm[lc[t]] + sm[rc[t]]
sz[t] = 1 + sz[lc[t]] + sz[rc[t]]
def push_down(t):
if rev[t]:
if lc[t] != NIL:
l2 = clone(lc[t]); lc[t] = l2
rev[l2] = not rev[l2]
if rc[t] != NIL:
r2 = clone(rc[t]); rc[t] = r2
rev[r2] = not rev[r2]
lc[t], rc[t] = rc[t], lc[t]
rev[t] = False
def split(t, k):
if t == NIL:
return NIL, NIL
t = clone(t)
push_down(t)
ls = sz[lc[t]]
if k <= ls:
a, b = split(lc[t], k)
lc[t] = b
pull(t)
return a, t
else:
a, b = split(rc[t], k - ls - 1)
rc[t] = a
pull(t)
return t, b
def merge(a, b):
if a == NIL:
return b
if b == NIL:
return a
if pr[a] > pr[b]:
a = clone(a)
push_down(a)
rc[a] = merge(rc[a], b)
pull(a)
return a
else:
b = clone(b)
push_down(b)
lc[b] = merge(a, lc[b])
pull(b)
return b
def build(arr):
root = NIL
for v in arr:
root = merge(root, new_node(v))
return root
def range_sum(t, l, r):
a, bc = split(t, l - 1)
b, c = split(bc, r - l + 1)
return sm[b]
def range_reverse(t, l, r):
a, bc = split(t, l - 1)
b, c = split(bc, r - l + 1)
if b != NIL:
b = clone(b)
rev[b] = not rev[b]
return merge(merge(a, b), c)
def swap_positions(t, x, y):
if x > y:
x, y = y, x
if x == y:
return t
a, bcd = split(t, x - 1)
b, cd = split(bcd, 1)
c, d = split(cd, y - x - 1)
e, f = split(d, 1)
if b != NIL and e != NIL:
b = clone(b); e = clone(e)
val[b], val[e] = val[e], val[b]
pull(b); pull(e)
return merge(merge(merge(merge(a, b), c), e), f)
versions = [build(A)]
out = []
for _ in range(Q):
t = input_data[idx]; idx += 1
if t == b"1":
v, l, r = int(input_data[idx]), int(input_data[idx+1]), int(input_data[idx+2]); idx += 3
versions.append(range_reverse(versions[v], l, r))
elif t == b"2":
v, l, r = int(input_data[idx]), int(input_data[idx+1]), int(input_data[idx+2]); idx += 3
out.append(str(range_sum(versions[v], l, r)))
elif t == b"3":
v, x, y = int(input_data[idx]), int(input_data[idx+1]), int(input_data[idx+2]); idx += 3
versions.append(swap_positions(versions[v], x, y))
else:
v = int(input_data[idx]); idx += 1
versions.append(versions[v])
print("\n".join(out))
main()
計算量: 各更新クエリは $O(\log N)$ 個のノードを新規複製し、時間・空間ともに1クエリあたり $O(\log N)$(ならし)。全体で $O((N+Q)\log N)$。ランダムクエリ200トライアル×30クエリで、素朴な「配列を毎回コピーする」参照実装と全出力が一致することを確認済み。
Step-by-Step 解説
1永続化の核心「コピー・オン・ライト」
書き換える直前に必ず
書き換える直前に必ず
clone(t)して新しいノード番号を作り、そこを書き換える。元のノードは変更されないので、古いバージョンの根から辿れば昔の状態がそのまま残る。2遅延反転フラグも「複製してから」立てる
push_downで子に伝播する際も子を複製してからフラグを立てる。複製し忘れると他のバージョンで共有されているノードまで書き換わってしまう。3
各クエリの後に新しい根ノードを
versions[]配列で世代を管理する各クエリの後に新しい根ノードを
versionsに追加する。クエリ2(区間和取得)だけは新バージョンを作らない。4値の交換は「1要素区間として取り出して戻す」
位置$x$と$y$をそれぞれ1要素の部分木として
位置$x$と$y$をそれぞれ1要素の部分木として
splitで切り出し、値だけ交換してからmergeで元の位置に戻す。5分岐(クエリ4)は
あるバージョンをそのまま新バージョンとして複製するのは根のポインタをコピーするだけで済む。
O(1)あるバージョンをそのまま新バージョンとして複製するのは根のポインタをコピーするだけで済む。
よくあるミス
| ミス | 原因 | 正しい書き方 |
|---|---|---|
push_downで子を複製せず直接書き換える | 通常のTreapのコードをそのまま流用してしまう | 子をcloneしてからフラグを立て、親のlc/rcを新しいノード番号に差し替える |
split/mergeの再帰で複製し忘れる箇所がある | 一部のパスだけ複製し、他は元ノードを書き換えてしまう | 再帰に入るたびに必ずcloneしてからpush_down・書き換えを行う |
versionsにクエリ2の結果も追加してしまう | 「クエリ後に必ずappendする」という思い込み | クエリ2は新バージョンを作らないのでversionsに追加しない |
| 優先度をクエリごとに再生成してしまう | random.random()をノード作成以外の場所で呼んでしまう | 優先度はノード作成時に1回だけ決め、以後は固定値として扱う |
次のステップ
- 発展: クエリに「バージョン$v_1$と$v_2$をマージする」を追加し、永続Treapの
mergeを使った履歴統合を実装する - 発展: Persistent Segment Tree(区間和のみ・反転なし)と比較し、反転操作の有無がデータ構造選択に与える影響を整理する
- 発展: 削除・挿入クエリを追加し、列の長さが動的に変わる場合の永続化を考える