結果

問題 No.3755 Root for Your Route
コンテスト
ユーザー marc2825
提出日時 2026-08-19 23:44:03
言語 PyPy3
(7.3.23 + ACL)
コンパイル:
pypy3 -mpy_compile _filename_
実行:
pypy3 _filename_
結果
AC  
実行時間 1,053 ms / 3,000 ms
+ 755µs
コード長 5,734 bytes
記録
記録タグの例:
初AC ショートコード 純ショートコード 純主流ショートコード 最速実行時間
コンパイル時間 66 ms
コンパイル使用メモリ 82,420 KB
実行使用メモリ 144,032 KB
最終ジャッジ日時 2026-10-02 21:03:15
合計ジャッジ時間 26,695 ms
ジャッジサーバーID
(参考情報)
judge2_0 / judge1_0
このコードへのチャレンジ
(要ログイン)
ファイルパターン 結果
sample AC * 1
other AC * 39
権限があれば一括ダウンロードができます

ソースコード

diff #
raw source code

import sys

def main():
    data = sys.stdin.buffer.read().split()
    N = int(data[0])
    A = [0] * (N + 1)
    for i in range(N):
        A[i + 1] = int(data[i + 1])

    M = N - 1
    head = [0] * (N + 2)
    us = [0] * M; vs = [0] * M
    p = N + 1
    for i in range(M):
        u = int(data[p]); v = int(data[p + 1]); p += 2
        us[i] = u; vs[i] = v
        head[u + 1] += 1; head[v + 1] += 1
    for i in range(1, N + 2):
        head[i] += head[i - 1]
    pos = head[:]
    adj = [0] * (2 * M)
    for i in range(M):
        u = us[i]; v = vs[i]
        adj[pos[u]] = v; pos[u] += 1
        adj[pos[v]] = u; pos[v] += 1

    NEG = -(1 << 62)
    removed = bytearray(N + 1)
    sz  = [0] * (N + 1)
    par = [0] * (N + 1)
    dep = [0] * (N + 1)
    Sm  = [0] * (N + 1)
    Xv  = [0] * (N + 1)
    bs  = [0] * (N + 1)
    g   = [NEG] * (N + 1)
    ans = [NEG] * (N + 1)
    order = [0] * (N + 1)
    comp  = [0] * (N + 1)
    bstart = [0] * (N + 2)

    size4 = 4 * (N + 2)
    lm = [0] * size4; lb = [0] * size4; lst = [-1] * size4
    cur = 0

    stack = [1]
    while stack:
        root = stack.pop()
        # ---- 連結成分の BFS と部分木サイズ ----
        par[root] = 0; order[0] = root; cnt = 1; i = 0
        while i < cnt:
            v = order[i]; i += 1; pv = par[v]
            for j in range(head[v], head[v + 1]):
                u = adj[j]
                if u != pv and not removed[u]:
                    par[u] = v; order[cnt] = u; cnt += 1
        tot = cnt
        for i in range(tot):
            sz[order[i]] = 1
        for i in range(tot - 1, 0, -1):
            v = order[i]; sz[par[v]] += sz[v]

        # ---- 重心を降りて探す ----
        c = root
        while True:
            nx = -1
            for j in range(head[c], head[c + 1]):
                u = adj[j]
                if u != par[c] and not removed[u] and sz[u] * 2 > tot:
                    nx = u; break
            if nx < 0: break
            c = nx

        # ---- 重心成分を枝ごとに展開し dep / Sum / X を計算 ----
        Ac = A[c]
        comp[0] = c; par[c] = 0; dep[c] = 0; Sm[c] = Ac; Xv[c] = Ac
        cnt = 1; nb = 0; maxd = 1
        for j0 in range(head[c], head[c + 1]):
            r0 = adj[j0]
            if removed[r0]: continue
            bstart[nb] = cnt; nb += 1
            par[r0] = c; dep[r0] = 1; Sm[r0] = Ac + A[r0]; Xv[r0] = Ac + A[r0] - 1
            comp[cnt] = r0; cnt += 1
            i = cnt - 1
            while i < cnt:
                v = comp[i]; i += 1
                pv = par[v]; dv = dep[v]; sv = Sm[v]
                for j in range(head[v], head[v + 1]):
                    u = adj[j]
                    if u != pv and not removed[u]:
                        par[u] = v; d = dv + 1; dep[u] = d
                        s2 = sv + A[u]; Sm[u] = s2
                        Xv[u] = s2 - d * (d + 1) // 2
                        comp[cnt] = u; cnt += 1
                        if d > maxd: maxd = d
        bstart[nb] = cnt

        maxX = NEG
        for i in range(cnt):
            v = comp[i]
            if Xv[v] > maxX: maxX = Xv[v]
            bs[v] = Ac                      # w = c(片腕のみ)の場合
        LCN = maxd

        # ---- prefix / suffix の 2 方向スキャン ----
        for direction in (0, 1):
            rng = range(nb) if direction == 0 else range(nb - 1, -1, -1)
            first = True
            cur += 1
            for k in rng:
                lo = bstart[k]; hi = bstart[k + 1]
                if not first:                        # クエリ
                    for i in range(lo, hi):
                        v = comp[i]; x = dep[v]
                        node = 1; l = 0; r = LCN; res = bs[v]
                        while True:
                            if lst[node] == cur:
                                t = lm[node] * x + lb[node]
                                if t > res: res = t
                            if l == r: break
                            mid = (l + r) >> 1
                            if x <= mid: node <<= 1; r = mid
                            else: node = node * 2 + 1; l = mid + 1
                        bs[v] = res
                first = False
                for i in range(lo, hi):              # 挿入
                    v = comp[i]; m = -dep[v]; b = Xv[v]
                    node = 1; l = 0; r = LCN
                    while True:
                        if lst[node] != cur:
                            lst[node] = cur; lm[node] = m; lb[node] = b
                            break
                        m2 = lm[node]; b2 = lb[node]
                        lef = (m * l + b) > (m2 * l + b2)
                        mid = (l + r) >> 1
                        if (m * mid + b) > (m2 * mid + b2):
                            lm[node] = m; lb[node] = b
                            m = m2; b = b2; mb = True
                        else:
                            mb = False
                        if l == r: break
                        if lef != mb: node <<= 1; r = mid
                        else: node = node * 2 + 1; l = mid + 1

        # ---- val(u) を求め、部分木 max で各頂点へ配る ----
        for i in range(1, cnt):
            v = comp[i]
            g[v] = Xv[v] + bs[v] - Ac
        g[c] = maxX
        for i in range(cnt - 1, 0, -1):
            v = comp[i]; pv = par[v]
            if g[v] > g[pv]: g[pv] = g[v]
        for i in range(cnt):
            v = comp[i]
            if g[v] > ans[v]: ans[v] = g[v]

        removed[c] = 1
        for j in range(head[c], head[c + 1]):
            u = adj[j]
            if not removed[u]: stack.append(u)

    res = min(ans[1:N + 1])
    sys.stdout.write(str(res) + "\n")

main()
0