Official

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\)より大きい要素の個数を数える」問題に帰着できる - 区間クエリを高速に処理するため、平方分割を用いて平衡二分探索木の代わりにソート済みバケットを管理する

アルゴリズム

  1. オイラーツアー: DFSを用いて各ノードの入時刻(in_time)と出時刻(out_time)を記録します。これにより、社員\(v\)の部下全体が配列上の連続区間\([in\_time[v]+1, out\_time[v]]\)に対応します。

  2. 平方分割: 配列を\(\sqrt{N}\)程度のサイズのバケットに分割し、各バケット内の値をソートして保持します。

  3. 更新処理: 給与更新時には、該当するバケットから古い値を削除し、新しい値を挿入します。バケットは常にソート状態を保ちます。

  4. クエリ処理: 区間クエリでは、完全に含まれるバケットについては二分探索で\(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: