E - 会社の組織と給与 / Company Organization and Salaries Editorial by admin
DeepSeek V3概要
木構造で表される会社組織において、各社員の給与を更新しつつ、指定された社員の部下の中でその社員の給与より大きい給与を持つ人数を求めるクエリを効率的に処理する問題です。
考察
問題の核となる操作は、以下の2点です: 1. 社員の給与の更新 2. ある社員の部下全員の中で、その社員の給与より大きい給与を持つ人数の計算
素朴なアプローチでは、クエリ2ごとに部下全員を走査して給与を比較すると、最悪の場合\(O(NQ)\)となり、制約条件\(N, Q \leq 10^5\)では間に合いません。
効率的に処理するためには、以下の観察が重要です: - 木をDFS順(オイラーツアー)で辿ると、各社員の部下全体が連続した区間として表せる - クエリ2は「区間\([l, r]\)内で値\(x\)より大きい要素の個数を数える」問題に帰着できる - 区間クエリを高速に処理するため、平方分割を用いて平衡二分探索木の代わりにソート済みバケットを管理する
アルゴリズム
オイラーツアー: DFSを用いて各ノードの入時刻(in_time)と出時刻(out_time)を記録します。これにより、社員\(v\)の部下全体が配列上の連続区間\([in\_time[v]+1, out\_time[v]]\)に対応します。
平方分割: 配列を\(\sqrt{N}\)程度のサイズのバケットに分割し、各バケット内の値をソートして保持します。
更新処理: 給与更新時には、該当するバケットから古い値を削除し、新しい値を挿入します。バケットは常にソート状態を保ちます。
クエリ処理: 区間クエリでは、完全に含まれるバケットについては二分探索で\(x\)より大きい要素数を数え、部分的な区間は直接走査します。
計算量
- 時間計算量: \(O(Q\sqrt{N}\log N)\)
- 更新処理: \(O(\sqrt{N}\log N)\)(バケットの削除・挿入)
- クエリ処理: \(O(\sqrt{N}\log N)\)(バケットの二分探索と部分走査)
- 空間計算量: \(O(N)\)
実装のポイント
オイラーツアーは非再帰DFSで実装し、スタックオーバーフローを防ぐ
バケットサイズは\(\sqrt{N}\)に設定し、計算量のバランスを取る
バケットへの挿入・削除には
bisectモジュールを使用して効率化部分区間の処理では境界ケース(空区間など)に注意する
ソースコード
import sys
import math
import bisect
def main():
data = sys.stdin.read().split()
if not data:
return
it = iter(data)
n = int(next(it)); q = int(next(it))
A = [0] * (n+1)
for i in range(1, n+1):
A[i] = int(next(it))
graph = [[] for _ in range(n+1)]
for i in range(2, n+1):
p = int(next(it))
graph[p].append(i)
in_time = [0] * (n+1)
out_time = [0] * (n+1)
time = 0
stack = [1]
while stack:
u = stack.pop()
if u > 0:
time += 1
in_time[u] = time
stack.append(-u)
for v in reversed(graph[u]):
stack.append(v)
else:
u = -u
out_time[u] = time
arr = [0] * (n+1)
for i in range(1, n+1):
arr[in_time[i]] = A[i]
size = int(math.isqrt(n)) + 1
bucket_count = (n + size) // size
buckets = [[] for _ in range(bucket_count)]
for i in range(1, n+1):
idx_b = (i-1) // size
buckets[idx_b].append(arr[i])
for i in range(bucket_count):
buckets[i].sort()
def update(pos, old_val, new_val):
idx_b = (pos-1) // size
arr[pos] = new_val
bucket = buckets[idx_b]
# Remove old_val
pos_in_bucket = bisect.bisect_left(bucket, old_val)
del bucket[pos_in_bucket]
# Insert new_val
bisect.insort(bucket, new_val)
def query(l, r, x):
res = 0
left_b = (l-1) // size
right_b = (r-1) // size
if left_b == right_b:
for i in range(l, r+1):
if arr[i] > x:
res += 1
else:
for i in range(l, min((left_b+1)*size+1, n+1)):
if arr[i] > x:
res += 1
for b in range(left_b+1, right_b):
bucket = buckets[b]
cnt = len(bucket) - bisect.bisect_right(bucket, x)
res += cnt
for i in range(right_b*size+1, r+1):
if arr[i] > x:
res += 1
return res
output_lines = []
for _ in range(q):
t = next(it)
if t == '1':
v = int(next(it)); x = int(next(it))
pos = in_time[v]
old_val = arr[pos]
update(pos, old_val, x)
else:
v = int(next(it))
l = in_time[v] + 1
r = out_time[v]
if l > r:
output_lines.append("0")
else:
x = arr[in_time[v]]
cnt = query(l, r, x)
output_lines.append(str(cnt))
sys.stdout.write("\n".join(output_lines))
if __name__ == "__main__":
main()
この解説は deepseekv3 によって生成されました。
posted:
last update: