結果

問題 No.3720 Balanced Reduction
コンテスト
ユーザー kidodesu
提出日時 2026-09-18 23:19:52
言語 PyPy3
(7.3.23 + ACL)
コンパイル:
pypy3 -mpy_compile _filename_
実行:
pypy3 _filename_
結果
WA  
実行時間 -
コード長 4,042 bytes
記録
記録タグの例:
初AC ショートコード 純ショートコード 純主流ショートコード 最速実行時間
コンパイル時間 77 ms
コンパイル使用メモリ 82,176 KB
実行使用メモリ 181,044 KB
最終ジャッジ日時 2026-09-18 23:20:05
合計ジャッジ時間 5,741 ms
ジャッジサーバーID
(参考情報)
judge1_0 / judge2_0
このコードへのチャレンジ
(要ログイン)
ファイルパターン 結果
sample AC * 4
other AC * 14 WA * 2
権限があれば一括ダウンロードができます

ソースコード

diff #
raw source code

def main():
    n, k = list(map(int, input().split()))
    A = list(map(int, input().split()))
    node = [[] for _ in range(n)]
    E = [0] * n
    for _ in range(n):
        u, v = list(map(lambda x: int(x)-1, input().split()))
        node[u].append(v)
        node[v].append(u)
        E[u] += 1
        E[v] += 1
    S = [u for u in range(n) if E[u] == 1]
    ans0 = 0
    def cal(z):
        return (z-1) // (2*k) + 1
    while S:
        u = S.pop()
        if A[u] == 0:
            pass
        elif A[u] < k:
            return -1
        ans0 += cal(A[u])
        for v in node[u]:
            if E[v]:
                E[v] -= 1
                E[u] -= 1
                A[v] -= A[u]
                A[u] = 0
                if E[v] == 1:
                    S.append(v)
    X = set()
    u = 0
    while not E[u]:
        u += 1
    B = []
    while 1:
        X.add(u)
        B.append(A[u])
        for v in node[u]:
            if not v in X and E[v]:
                u = v
                break
        else:
            break
    ans1 = 0
    N = len(B)
    A = B
    if N % 2:
        t0 = t1 = 0
        for u in range(N):
            if not u % 2:
                t0 += A[u]
            else:
                t1 += A[u]
        if (t0-t1) % 2: return -1
        x = (t0 - t1) // 2
        if x < 0: return -1
        elif 0 < x < k: return -1
        ans1 += cal(x)
        A[0] -= x
        A[-1] -= x
        for i in range(N-1):
            if A[i] < 0 or 0 < A[i] < k: return -1
            ans1 += cal(A[i])
            A[i+1] -= A[i]
    else:
        t0 = t1 = 0
        for u in range(N):
            if not u % 2:
                t0 += A[u]
            else:
                t1 += A[u]
        if t0 != t1: return -1
        X0 = []
        X1 = [0]
        for u in range(N-1):
            if not u % 2:
                X0.append(A[u])
            else:
                X1.append(A[u])
            A[u+1] -= A[u]
            A[u] = 0
        X0.sort()
        X1.sort()
        xx = -X1[0]
        X0 = [x0-xx for x0 in X0]
        X1 = [x1+xx for x1 in X1]
        #print(X0, X1)
        inf = 1<<60
        ans1 = inf
        ans2 = ans3 = 0
        for x in X0:
            if x < 0 or 0 < x < k:
                break
            else:
                ans2 += cal(x)
        else:
            for x in X1:
                if x < 0 or 0 < x < k:
                    break
                else:
                    ans2 += cal(x)
            else:
                ans1 = min(ans1, ans2)
        for x in X0:
            x -= X0[0]
            if x < 0 or 0 < x < k:
                break
            else:
                ans3 += cal(x)
        else:
            for x in X1:
                x += X0[0]
                if x < 0 or 0 < x < k:
                    break
                else:
                    ans3 += cal(x)
            else:
                ans1 = min(ans1, ans3)
        X0 = [x0-k for x0 in X0]
        X1 = [x1+k for x1 in X1]
        r = X0[0]-k
        t = 0
        if r < 0:
            pass
        else:
            for x in X0+X1:
                if x < 0 or 0 < x < k:
                    break
                else:
                    t += cal(x)
            else:
                F = []
                for x in X0:
                    for ki in [x//(2*k)*2*k, x//(2*k)*2*k-2*k]:
                        if ki < x and x-ki <= r:
                            F.append(((x-ki)*3-1))
                for x in X1:
                    x -= 1
                    for ki in [(x+2+2*k-1)//(2*k)*2*k, (x+2+2*k-1)//(2*k)*2*k+2*k]:
                        if x < ki and ki-x <= r:
                            F.append(((ki-x)*3+1))
                ans1 = min(ans1, t)
                F.sort()
                for s in F:
                    s %= 3
                    if s == 2:
                        t -= 1
                    else:
                        t += 1
                    ans1 = min(ans1, t)
        if 1<<59 <= ans1:
            return -1
    return ans0+ans1

print(main())
0