問題
$N$ 頂点の根付き木(根: 1)があり、各頂点 $v$ には値 $a_v$ がある。以下の $Q$ 個のクエリを処理せよ:
1 v x: 頂点 $v$ の値を $x$ に変更2 v: 頂点 $v$ の部分木内の全頂点の値の総和を答えよ3 u v: 頂点 $u$ から $v$ へのパス上の全頂点の値の最大値を答えよ
入力形式
N Q
a_1 a_2 ... a_N
p_2 p_3 ... p_N
クエリ1
...
制約
$2 \leq N \leq 2 \times 10^5$
$1 \leq Q \leq 2 \times 10^5$
$-10^9 \leq a_v, x \leq 10^9$
入出力例
入力例 1
7 5
1 2 3 4 5 6 7
1 1 2 2 3 3
2 1
3 4 6
1 3 10
2 1
3 4 6
出力例 1
28
6
31
10
ヒント (段階的開示)
ヒント1: 方向性
Euler Tour(DFS 順でのインタイム/アウトタイム付け)で部分木クエリを区間クエリに変換。HLD でパスクエリを $O(\log N)$ 個の区間クエリに変換。どちらもセグメント木 1 本で処理できる。
ヒント2: アプローチ
HLD:
subtree_size を DFS で計算、各ノードの「重い子」= 最大部分木サイズの子、重いパスに沿って DFS 順でノード番号付け。部分木クエリ: in[v] から in[v] + subtree_size[v] - 1。パスクエリ: 重いパスを登りながら区間クエリ。ヒント3: 誘導
def dfs_hld(v, par, depth):
sz[v] = 1
heavy[v] = -1
for c in children[v]:
if c == par: continue
dfs_hld(c, v, depth + 1)
sz[v] += sz[c]
if heavy[v] == -1 or sz[c] > sz[heavy[v]]:
heavy[v] = c
def decompose(v, h, par):
pos[v] = timer[0]; timer[0] += 1
head[v] = h
if heavy[v] != -1:
decompose(heavy[v], h, v)
for c in children[v]:
if c != par and c != heavy[v]:
decompose(c, c, v)
模範解答 (Python)
import sys
from sys import setrecursionlimit
input = sys.stdin.readline
def solve():
N, Q = map(int, input().split())
a = list(map(int, input().split()))
parents = list(map(int, input().split()))
children = [[] for _ in range(N + 1)]
for i in range(2, N + 1):
p = parents[i - 2]
children[p].append(i)
children[i].append(p)
sz = [1] * (N + 1)
depth = [0] * (N + 1)
par = [0] * (N + 1)
heavy = [-1] * (N + 1)
# DFS1: subtree_size と heavy child
order = []
stack = [(1, 0, False)]
while stack:
v, p, visited = stack.pop()
if visited:
best = -1
for c in children[v]:
if c == p:
continue
sz[v] += sz[c]
if best == -1 or sz[c] > sz[best]:
best = c
heavy[v] = best
else:
order.append((v, p))
par[v] = p
stack.append((v, p, True))
for c in children[v]:
if c != p:
depth[c] = depth[v] + 1
stack.append((c, v, False))
# DFS2: pos と head
pos = [0] * (N + 1)
head = [0] * (N + 1)
in_time = [0] * (N + 1)
out_time = [0] * (N + 1)
timer = [0]
stack2 = [(1, 0, 1)]
while stack2:
v, p, h = stack2.pop()
pos[v] = timer[0]
in_time[v] = timer[0]
head[v] = h
timer[0] += 1
out_time[v] = in_time[v] + sz[v] - 1
for c in children[v]:
if c != p and c != heavy[v]:
stack2.append((c, v, c))
if heavy[v] != -1:
stack2.append((heavy[v], v, h))
SIZE = N
seg_sum = [0] * (4 * SIZE)
seg_max = [-10**18] * (4 * SIZE)
def build(node, l, r, arr):
if l == r:
seg_sum[node] = arr[l]
seg_max[node] = arr[l]
return
m = (l + r) // 2
build(2*node, l, m, arr)
build(2*node+1, m+1, r, arr)
seg_sum[node] = seg_sum[2*node] + seg_sum[2*node+1]
seg_max[node] = max(seg_max[2*node], seg_max[2*node+1])
def update(node, l, r, idx, val):
if l == r:
seg_sum[node] = val
seg_max[node] = val
return
m = (l + r) // 2
if idx <= m:
update(2*node, l, m, idx, val)
else:
update(2*node+1, m+1, r, idx, val)
seg_sum[node] = seg_sum[2*node] + seg_sum[2*node+1]
seg_max[node] = max(seg_max[2*node], seg_max[2*node+1])
def query_sum(node, l, r, ql, qr):
if qr < l or r < ql:
return 0
if ql <= l and r <= qr:
return seg_sum[node]
m = (l + r) // 2
return query_sum(2*node, l, m, ql, qr) + query_sum(2*node+1, m+1, r, ql, qr)
def query_max(node, l, r, ql, qr):
if qr < l or r < ql:
return -10**18
if ql <= l and r <= qr:
return seg_max[node]
m = (l + r) // 2
return max(query_max(2*node, l, m, ql, qr), query_max(2*node+1, m+1, r, ql, qr))
init_arr = [0] * N
for v in range(1, N + 1):
init_arr[pos[v]] = a[v - 1]
build(1, 0, N - 1, init_arr)
def path_max(u, v):
res = -10**18
while head[u] != head[v]:
if depth[head[u]] < depth[head[v]]:
u, v = v, u
res = max(res, query_max(1, 0, N-1, pos[head[u]], pos[u]))
u = par[head[u]]
if depth[u] > depth[v]:
u, v = v, u
res = max(res, query_max(1, 0, N-1, pos[u], pos[v]))
return res
out = []
for _ in range(Q):
line = list(map(int, input().split()))
if line[0] == 1:
v, x = line[1], line[2]
update(1, 0, N-1, pos[v], x)
elif line[0] == 2:
v = line[1]
out.append(query_sum(1, 0, N-1, in_time[v], out_time[v]))
else:
u, v = line[1], line[2]
out.append(path_max(u, v))
print('\n'.join(map(str, out)))
solve()
Step-by-Step 解説
1HLD の基本概念
Heavy-Light Decomposition は木のパスを $O(\log N)$ 個の重いパス上の連続区間に分解する技法。重い辺: 親から最大部分木サイズの子への辺。性質: 根から任意の頂点へのパスで軽い辺を通る回数は $O(\log N)$ 以下。
Heavy-Light Decomposition は木のパスを $O(\log N)$ 個の重いパス上の連続区間に分解する技法。重い辺: 親から最大部分木サイズの子への辺。性質: 根から任意の頂点へのパスで軽い辺を通る回数は $O(\log N)$ 以下。
2DFS 順の番号付け
重い子を先に訪問することで、同一重いパス上のノードが連続した番号を持つ。結果として重いパスが連続区間に対応。
重い子を先に訪問することで、同一重いパス上のノードが連続した番号を持つ。結果として重いパスが連続区間に対応。
3部分木クエリ
in_time[v] から out_time[v] = in_time[v] + sz[v] - 1 が部分木の範囲。セグメント木の区間和クエリで $O(\log N)$。4パスクエリ(HLD)
$u$ と $v$ の LCA を求めながら上に登る。
$u$ と $v$ の LCA を求めながら上に登る。
head[u] != head[v] なら深い方の head から該当ノードまで区間クエリし親の親へ移動。同じ head になったら浅い方から深い方まで区間クエリ。よくあるミス
| ミス | 原因 | 正しい書き方 |
|---|---|---|
部分木の範囲計算が sz[v] を足しすぎる | out = in + sz - 1 が正しい | out_time[v] = in_time[v] + sz[v] - 1 |
| path_max で浅い/深いの判定ミス | head の深さで比較 | depth[head[u]] < depth[head[v]] |
| 再帰 DFS でスタックオーバーフロー | N = 2×10^5 で再帰制限超え | 反復 DFS を実装する |
| sum と max を別セグ木で管理 | 更新が2回必要になりミスしやすい | 1つのセグ木に両方保持 |
次のステップ
- 発展問題: 辺に重みがある場合のパス上の最大辺重みを求めよ(各辺の重みを子ノードに持たせて HLD を適用)