結果

問題 No.3749 Three Jugs
コンテスト
ユーザー Naru820
提出日時 2026-09-16 20:56:13
言語 PyPy3
(7.3.23 + ACL)
コンパイル:
pypy3 -mpy_compile _filename_
実行:
pypy3 _filename_
結果
TLE  
実行時間 -
コード長 6,226 bytes
記録
記録タグの例:
初AC ショートコード 純ショートコード 純主流ショートコード 最速実行時間
コンパイル時間 974 ms
コンパイル使用メモリ 82,176 KB
実行使用メモリ 156,156 KB
最終ジャッジ日時 2026-09-25 20:55:10
合計ジャッジ時間 74,400 ms
ジャッジサーバーID
(参考情報)
judge2_0 / judge4_0
このコードへのチャレンジ
(要ログイン)
ファイルパターン 結果
sample AC * 1
other AC * 22 TLE * 1
権限があれば一括ダウンロードができます

ソースコード

diff #
raw source code

import sys

INF = 1 << 62
PAIRS = (
    (0, 1, 2),
    (0, 2, 1),
    (1, 0, 2),
    (1, 2, 0),
    (2, 0, 1),
    (2, 1, 0),
)


def solve(A, B, C):
    A0, A1, A2 = A
    W = B[0] + B[1] + B[2]

    def boundary(x):
        x0, x1, x2 = x
        return (
            (x0 == 0 or x0 == A0)
            + (x1 == 0 or x1 == A1)
            + (x2 == 0 or x2 == A2)
        )

    if B == C:
        return 0

    if boundary(C) == 0:
        return -1

    start = []

    for i, j, k in PAIRS:
        x = list(B)
        d = min(x[i], A[j] - x[j])
        x[i] -= d
        x[j] += d
        x = tuple(x)

        if x == C:
            return 1

        start.append(x)

    if boundary(C) >= 2:
        return 2

    start_set = set(start)

    # id ごとの辺グループ
    length = []
    weight = []
    initial = []

    order0 = []
    order1 = []

    for i, j, k in PAIRS:
        Ai = A[i]
        Aj = A[j]
        Ak = A[k]

        lo = max(0, W - Ai - Aj + 1)
        hi = min(Ak, W - 1)

        if lo > hi:
            continue

        cuts = [lo, hi + 1]

        # 頂点・印・目標・遷移の種類が変化する位置
        for z in (
            0,
            Ak,
            W - Ai,
            W - Aj,
            B[k],
            C[k],
        ):
            if lo <= z <= hi:
                cuts.append(z)
                cuts.append(z + 1)

        for x in start:
            z = x[k]
            if lo <= z <= hi:
                cuts.append(z)
                cuts.append(z + 1)

        cuts = sorted(set(cuts))

        for p in range(len(cuts) - 1):
            left = cuts[p]
            right = cuts[p + 1]
            seg_len = right - left

            for r in (0, 1):
                # r=0: 通常順, r=1: 逆順
                z = right - 1 if r else left
                sign = -1 if r else 1

                u = [0, 0, 0]
                v = [0, 0, 0]

                u[k] = v[k] = z

                s = W - z

                ui = min(Ai, s)
                u[i] = ui
                u[j] = s - ui

                vj = min(Aj, s)
                v[j] = vj
                v[i] = s - vj

                ut = tuple(u)
                vt = tuple(v)

                # boundary(u), boundary(v) をここで一度だけ計算
                bu = (
                    (u[0] == 0 or u[0] == A0)
                    + (u[1] == 0 or u[1] == A1)
                    + (u[2] == 0 or u[2] == A2)
                )

                bv = (
                    (v[0] == 0 or v[0] == A0)
                    + (v[1] == 0 or v[1] == A1)
                    + (v[2] == 0 or v[2] == A2)
                )

                if ut == C:
                    mu = -1
                elif ut == B:
                    mu = 0
                elif ut in start_set:
                    mu = 1
                else:
                    mu = 2 if bu >= 2 else 3

                if vt == C:
                    mv = -1
                elif vt == B:
                    mv = 0
                elif vt in start_set:
                    mv = 1
                else:
                    mv = 2 if bv >= 2 else 3

                # v から次に行く操作
                ni = j
                nj = i

                if bv == 1:
                    if v[i] == 0:
                        ni = k
                        nj = i
                    else:
                        ni = j
                        nj = k

                idx = len(length)

                # 始点側
                order0.append((
                    mu == 3,
                    9 * r + 3 * i + j,
                    sign * z,
                    idx,
                ))

                # 終点側
                order1.append((
                    mv == 3,
                    9 * (1 - r) + 3 * ni + nj,
                    -sign * v[3 - ni - nj],
                    idx,
                ))

                length.append(seg_len)
                weight.append(1)
                initial.append(mu)

    order0.sort()
    order1.sort()

    row0 = [x[3] for x in order0]
    row1 = [x[3] for x in order1]

    dist = []

    # 特殊状態は先頭に連続している
    for idx in row0:
        m = initial[idx]
        if m == 3:
            break
        dist.append(m)

    keep = len(dist)
    L = length
    WGT = weight

    while len(row0) > keep:
        a = row0[-1]
        b = row1[-1]

        if a == b:
            row0.pop()
            row1.pop()
            continue

        # 本数が多い方を勝者にする
        if L[a] >= L[b]:
            side = 0
            win = a
            other = row1
        else:
            side = 1
            win = b
            other = row0

        pos = other.index(win)

        # 勝者より後ろにあるグループの本数の合計
        total = 0
        for t in range(pos + 1, len(other)):
            total += L[other[t]]

        q = L[win] // total

        if q:
            # 一巡分をまとめて処理
            L[win] %= total

            add = q * WGT[win]

            for t in range(pos + 1, len(other)):
                WGT[other[t]] += add

        else:
            lose = other[-1]

            L[win] -= L[lose]
            WGT[lose] += WGT[win]

            # lose を win の直後へ移す
            other.pop()
            other.insert(pos + 1, lose)

        if L[win] == 0:
            if side == 0:
                row0.pop()
            else:
                row1.pop()

            del other[pos]

    ans = INF

    for i, d in enumerate(dist):
        if d != -1:
            continue

        idx = row1[i]
        p = row0.index(idx)

        if dist[p] >= 0:
            cand = WGT[idx] + dist[p]
            if cand < ans:
                ans = cand

    return -1 if ans == INF else ans


def main():
    data = list(map(int, sys.stdin.buffer.read().split()))
    it = iter(data)

    T = next(it)
    out = []

    for _ in range(T):
        A = (next(it), next(it), next(it))
        B = (next(it), next(it), next(it))
        C = (next(it), next(it), next(it))

        out.append(str(solve(A, B, C)))

    sys.stdout.write("\n".join(out))


if __name__ == "__main__":
    main()
0