結果

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

ソースコード

diff #
raw source code

NEG = -10**30


# 直線 y = ax + b の追加、最大値取得
class LiChao:
    def __init__(self, n):
        self.n = n

        size = 4 * (n + 1)
        self.a = [0] * size
        self.b = [0] * size
        self.used = [0] * size
        self.ver = 0

    # 中身を実際には消さず、世代だけ進める
    def clear(self):
        self.ver += 1

    def add(self, a, b):
        A = self.a
        B = self.b
        used = self.used
        ver = self.ver

        k, l, r = 1, 0, self.n

        while True:
            if used[k] != ver:
                used[k] = ver
                A[k] = a
                B[k] = b
                return

            c = A[k]
            d = B[k]
            m = (l + r) // 2

            left = a * l + b > c * l + d
            mid = a * m + b > c * m + d

            if mid:
                A[k], a = a, c
                B[k], b = b, d

            if l == r:
                return

            if left != mid:
                k *= 2
                r = m
            else:
                k = k * 2 + 1
                l = m + 1

    def query(self, x):
        A = self.a
        B = self.b
        used = self.used
        ver = self.ver

        k, l, r = 1, 0, self.n
        res = NEG

        while True:
            if used[k] == ver:
                val = A[k] * x + B[k]
                if val > res:
                    res = val

            if l == r:
                return res

            m = (l + r) // 2

            if x <= m:
                k *= 2
                r = m
            else:
                k = k * 2 + 1
                l = m + 1


N = int(input())
A = list(map(int, input().split()))

G = [[] for _ in range(N)]

for _ in range(N - 1):
    u, v = map(int, input().split())
    u -= 1
    v -= 1

    G[u].append(v)
    G[v].append(u)


size = [0] * N
parent = [-1] * N
used = [False] * N

ans = A[:]
best = [NEG] * N
dp = [NEG] * N

# 全重心で使い回す
cht = LiChao(N)


def get_centroid(s):
    order = [s]
    parent[s] = -1

    for v in order:
        for u in G[v]:
            if used[u] or u == parent[v]:
                continue

            parent[u] = v
            order.append(u)

    for v in reversed(order):
        size[v] = 1

        for u in G[v]:
            if not used[u] and parent[u] == v:
                size[v] += size[u]

    n = len(order)

    for v in order:
        mx = n - size[v]

        for u in G[v]:
            if not used[u] and parent[u] == v and size[u] > mx:
                mx = size[u]

        if mx * 2 <= n:
            return v


# c を除いた1つの連結成分を集める
# B_v = sum(c -> v) - d_v(d_v+1)/2
def collect(c, root):
    comp = []
    stack = [(root, c, 1, A[c] + A[root] - 1)]

    while stack:
        v, p, d, b = stack.pop()
        comp.append((v, p, d, b))

        nd = d + 1

        for u in G[v]:
            if used[u] or u == p:
                continue

            # B_child = B_parent + A[child] - depth_child
            stack.append((u, v, nd, b + A[u] - nd))

    return comp


def solve(s):
    c = get_centroid(s)

    comps = []

    for u in G[c]:
        if not used[u]:
            comps.append(collect(c, u))

    if comps:
        # 左 -> 右
        cht.clear()
        cht.add(0, A[c])

        for comp in comps:
            for v, p, d, b in comp:
                best[v] = cht.query(d)

            for v, p, d, b in comp:
                cht.add(-d, b)

        # 右 -> 左
        cht.clear()
        cht.add(0, A[c])

        for comp in reversed(comps):
            for v, p, d, b in comp:
                q = cht.query(d)

                if q > best[v]:
                    best[v] = q

            for v, p, d, b in comp:
                cht.add(-d, b)

        best_c = A[c]

        for comp in comps:
            for v, p, d, b in comp:
                dp[v] = b - A[c] + best[v]

                if dp[v] > best_c:
                    best_c = dp[v]

            # 部分木最大値を子から親へ伝播
            for v, p, d, b in reversed(comp):
                val = dp[v]

                if val > ans[v]:
                    ans[v] = val

                if p != c and val > dp[p]:
                    dp[p] = val

        if best_c > ans[c]:
            ans[c] = best_c

    used[c] = True

    for u in G[c]:
        if not used[u]:
            solve(u)


solve(0)

print(min(ans))
0