E - 会社の組織と給与 / Company Organization and Salaries 解説 by admin
GPT 5.2 High概要
上司・部下関係が木で与えられるとき、給与の更新(1点更新)と「ある社員の部下のうち、自分より給与が高い人数」(部分木クエリ)を高速に処理します。
考察
重要な観察 1:部下集合は「部分木」
社員 \(v\) の部下は、根付き木でいう \(v\) の部分木(ただし \(v\) 自身を除く)に一致します。
ここで DFS の訪問順(Euler Tour)で各頂点に入った時刻 \(tin[v]\) を割り当てると、頂点 \(v\) の部分木に含まれる頂点は 配列上で連続区間 になります:
- 部分木(\(v\) 自身を含む) → \([tin[v],\ tout[v]]\)
- 部下(\(v\) 自身を除く) → \([tin[v]+1,\ tout[v]]\)
よって問題は次の形に変換できます:
配列上の区間 \([L,R]\) に含まれる値のうち、しきい値 \(T\) より大きい要素数を求める(しかも値は更新される)
素朴解が難しい理由
- クエリ2のたびに部下を全部なめて数えると最悪 \(O(N)\)、合計で \(O(NQ)\) になり \(10^{10}\) 規模で TLE。
- 「区間内で \(>T\) の個数」を高速に数えるには、値の大小を扱うデータ構造が必要。
- さらに更新クエリがあるため、静的なソートや累積和だけでは対応できません。
解決方針
区間クエリ(位置)と大小比較(値)を同時に扱うために、典型的な手法である
- Fenwick Tree(BIT)上に、さらに Fenwick Tree を持つ(2次元BIT)
を使います。
ただし値 \(A_i\) は最大 \(10^9\) と大きいので、各BITノードごとに必要な値だけを 座標圧縮 します(このためにクエリを先読みして、出現しうる値を集める「オフライン処理」を行います)。
アルゴリズム
1. Euler Tour で木を配列に潰す
反復DFSで \(tin[v], tout[v]\) を計算します(再帰は深さが \(10^5\) になり得るので避ける)。
- 訪問時に
timer += 1,tin[v] = timer - 部分木探索が終わった時点の
timerをtout[v]とする
すると、部下に対応する区間は: - \(L = tin[v] + 1\) - \(R = tout[v]\)
2. 「区間内で \(>T\)」を「\(\le T\)」に言い換える
区間内の要素数を total = R-L+1 とすると
- \(>T\) の個数 =
total - (<=T の個数)
よって必要なのは「区間内で \(\le T\) の個数」。
3. 2次元BIT(位置×値)
外側:位置(Euler Tour の添字)に対するBIT
add(pos, val, delta):位置posの値valをdeltaだけ増減(更新は+1と-1の組)sum_prefix(pos, val):位置 \(1..pos\) の中で、値 \(\le val\) の個数
区間 \([L,R]\) の「\(\le T\)」は
- sum_prefix(R, T) - sum_prefix(L-1, T)
内側:各外側BITノードに「値に対するBIT」
外側BITの各ノード k は、ある位置集合(BITの性質で決まる区間)を担当します。
そのノードが担当する位置に入りうる値だけを集めてソート・重複排除し、そこで座標圧縮したBIT(配列)を持たせます。
4. 座標圧縮を成立させるための「値の先読み」
更新により将来入る値も必要なので、次を行います:
- 各位置
tin[v]について- 初期値
A[v] - その社員が更新で取りうる値
xを収集しておく
- 初期値
そして、外側BITの構築時に
- 位置 p の候補値リストを、BITの伝播先 k = p, p+(p&-p), ... に配る
(コード中の coords[k].extend(vals))
これで各 coords[k] が「そのノードで現れうる値集合」になり、内部BITが作れます。
5. クエリ処理
- クエリ1
1 v x:- 位置
p = tin[v]の「古い給与」を-1、新しい給与を+1で反映
- 位置
- クエリ2
2 v:- \(L=tin[v]+1,\ R=tout[v]\)
- 部下がいない(\(L>R\))なら 0
- そうでなければ
le = count of <= A[v] in [L,R]ans = (R-L+1) - le
計算量
- 時間計算量: \(O\big((N+Q)\log^2 N\big)\)
- 更新・prefix集計が外側BITで \(O(\log N)\)、内側BITでも \(O(\log N)\) のため
- 空間計算量: \(O\big((N+Q)\log N\big)\)
- 各更新値が外側BITの複数ノードに配られるため(典型的に \(\log N\) 倍)
実装のポイント
部下は自分を含まないので、必ず区間を \([tin[v]+1, tout[v]]\) にする(ここがバグりやすい)。
Python では再帰DFSが危険なので、コードのように スタックで反復DFS にする。
2次元BITはメモリが重くなりがちなので、コードでは
内側BITを
array('i')(C配列相当)で持つcoords[k]をソート後にその場で重複削除 してメモリを節約している。ソースコード
import sys
from bisect import bisect_left, bisect_right
from array import array
def main():
data = list(map(int, sys.stdin.buffer.read().split()))
it = iter(data)
N = next(it)
Q = next(it)
A = [0] * (N + 1)
for i in range(1, N + 1):
A[i] = next(it)
children = [[] for _ in range(N + 1)]
if N >= 2:
for v in range(2, N + 1):
p = next(it)
children[p].append(v)
tin = [0] * (N + 1)
tout = [0] * (N + 1)
timer = 0
stack = [(1, 0)]
while stack:
v, st = stack.pop()
if st == 0:
timer += 1
tin[v] = timer
stack.append((v, 1))
ch = children[v]
for c in reversed(ch):
stack.append((c, 0))
else:
tout[v] = timer
queries = []
pos_vals = [[] for _ in range(N + 1)] # by Euler position
for v in range(1, N + 1):
pos_vals[tin[v]].append(A[v])
for _ in range(Q):
t = next(it)
if t == 1:
v = next(it)
x = next(it)
queries.append((1, v, x))
pos_vals[tin[v]].append(x)
else:
v = next(it)
queries.append((2, v))
coords = [[] for _ in range(N + 1)]
for p in range(1, N + 1):
vals = pos_vals[p]
if len(vals) > 1:
vals = list(set(vals))
vals.sort()
k = p
while k <= N:
coords[k].extend(vals)
k += k & -k
bits = [None] * (N + 1)
for k in range(1, N + 1):
arr = coords[k]
arr.sort()
m = 1
for i in range(1, len(arr)):
if arr[i] != arr[m - 1]:
arr[m] = arr[i]
m += 1
del arr[m:]
bits[k] = array('i', [0]) * (len(arr) + 1)
def add(pos, val, delta):
k = pos
while k <= N:
ck = coords[k]
idx = bisect_left(ck, val) + 1
bk = bits[k]
j = idx
blen = len(bk)
while j < blen:
bk[j] += delta
j += j & -j
k += k & -k
def sum_prefix(pos, val):
res = 0
k = pos
while k > 0:
ck = coords[k]
idx = bisect_right(ck, val)
bk = bits[k]
j = idx
while j > 0:
res += bk[j]
j -= j & -j
k -= k & -k
return res
for v in range(1, N + 1):
add(tin[v], A[v], 1)
out_lines = []
for q in queries:
if q[0] == 1:
_, v, x = q
old = A[v]
if old != x:
p = tin[v]
add(p, old, -1)
add(p, x, 1)
A[v] = x
else:
_, v = q
l = tin[v] + 1
r = tout[v]
if l > r:
out_lines.append("0")
continue
th = A[v]
total = r - l + 1
le = sum_prefix(r, th) - sum_prefix(l - 1, th)
out_lines.append(str(total - le))
sys.stdout.write("\n".join(out_lines))
if __name__ == "__main__":
main()
この解説は gpt-5.2-high によって生成されました。
投稿日時:
最終更新: