公式

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
  • 部分木探索が終わった時点の timertout[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 の値 valdelta だけ増減(更新は +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 によって生成されました。

投稿日時:
最終更新: