結果

問題 No.1320 Two Type Min Cost Cycle
コンテスト
ユーザー norioc
提出日時 2026-08-23 17:15:09
言語 PyPy3
(7.3.23)
コンパイル:
pypy3 -mpy_compile _filename_
実行:
pypy3 _filename_
結果
WA  
実行時間 -
コード長 9,455 bytes
記録
記録タグの例:
初AC ショートコード 純ショートコード 純主流ショートコード 最速実行時間
コンパイル時間 260 ms
コンパイル使用メモリ 95,976 KB
実行使用メモリ 97,280 KB
最終ジャッジ日時 2026-08-23 17:15:26
合計ジャッジ時間 14,252 ms
ジャッジサーバーID
(参考情報)
judge1_0 / judge2_0
このコードへのチャレンジ
(要ログイン)
ファイルパターン 結果
sample AC * 2 WA * 1
other AC * 18 WA * 28 RE * 11
権限があれば一括ダウンロードができます

ソースコード

diff #
raw source code

import typing
import sys
sys.setrecursionlimit(10**6)


def _ceil_pow2(n: int) -> int:
    x = 0
    while (1 << x) < n:
        x += 1

    return x


class SegTree:
    def __init__(self,
                 op: typing.Callable[[typing.Any, typing.Any], typing.Any],
                 e: typing.Any,
                 v: typing.Union[int, typing.List[typing.Any]]) -> None:
        self._op = op
        self._e = e

        if isinstance(v, int):
            v = [e] * v

        self._n = len(v)
        self._log = _ceil_pow2(self._n)
        self._size = 1 << self._log
        self._d = [e] * (2 * self._size)

        for i in range(self._n):
            self._d[self._size + i] = v[i]
        for i in range(self._size - 1, 0, -1):
            self._update(i)

    def set(self, p: int, x: typing.Any) -> None:
        assert 0 <= p < self._n

        p += self._size
        self._d[p] = x
        for i in range(1, self._log + 1):
            self._update(p >> i)

    def get(self, p: int) -> typing.Any:
        assert 0 <= p < self._n

        return self._d[p + self._size]

    def prod(self, left: int, right: int) -> typing.Any:
        assert 0 <= left <= right <= self._n
        sml = self._e
        smr = self._e
        left += self._size
        right += self._size

        while left < right:
            if left & 1:
                sml = self._op(sml, self._d[left])
                left += 1
            if right & 1:
                right -= 1
                smr = self._op(self._d[right], smr)
            left >>= 1
            right >>= 1

        return self._op(sml, smr)

    def all_prod(self) -> typing.Any:
        return self._d[1]

    def max_right(self, left: int,
                  f: typing.Callable[[typing.Any], bool]) -> int:
        assert 0 <= left <= self._n
        assert f(self._e)

        if left == self._n:
            return self._n

        left += self._size
        sm = self._e

        first = True
        while first or (left & -left) != left:
            first = False
            while left % 2 == 0:
                left >>= 1
            if not f(self._op(sm, self._d[left])):
                while left < self._size:
                    left *= 2
                    if f(self._op(sm, self._d[left])):
                        sm = self._op(sm, self._d[left])
                        left += 1
                return left - self._size
            sm = self._op(sm, self._d[left])
            left += 1

        return self._n

    def min_left(self, right: int,
                 f: typing.Callable[[typing.Any], bool]) -> int:
        assert 0 <= right <= self._n
        assert f(self._e)

        if right == 0:
            return 0

        right += self._size
        sm = self._e

        first = True
        while first or (right & -right) != right:
            first = False
            right -= 1
            while right > 1 and right % 2:
                right >>= 1
            if not f(self._op(self._d[right], sm)):
                while right < self._size:
                    right = 2 * right + 1
                    if f(self._op(self._d[right], sm)):
                        sm = self._op(self._d[right], sm)
                        right -= 1
                return right + 1 - self._size
            sm = self._op(self._d[right], sm)

        return 0

    def _update(self, k: int) -> None:
        self._d[k] = self._op(self._d[2 * k], self._d[2 * k + 1])


class UnionFind:
    def __init__(self, n: int):
        self.data = [-1] * (n+1)
        self.nexts = [i for i in range(n+1)]

    def root(self, a: int) -> int:
        if self.data[a] < 0: return a
        self.data[a] = self.root(self.data[a])
        return self.data[a]

    def unite(self, a: int, b: int) -> bool:
        pa = self.root(a)
        pb = self.root(b)
        if pa == pb: return False
        if self.data[pa] > self.data[pb]:
            pa, pb = pb, pa
        self.data[pa] += self.data[pb] # pa を pb をつなげる
        self.data[pb] = pa
        self.nexts[pa], self.nexts[pb] = self.nexts[pb], self.nexts[pa]
        return True

    def is_same(self, a: int, b: int) -> bool:
        return self.root(a) == self.root(b)

    def size(self, a: int) -> int:
        """a が属する集合のサイズ"""
        return -self.data[self.root(a)]

    def group(self, a: int):
        """a が属する集合"""
        yield a
        x = a
        while self.nexts[x] != a:
            x = self.nexts[x]
            yield x


class LCA:
    def __init__(self, n: int, adj: dict, root=0):
        """
        n: 頂点数
        adj: { 頂点: [隣接頂点, ...] }
        root: 根
        """
        sz = n.bit_length()
        parents = [[-1] * n for _ in range(sz)]
        dists = [-1] * n

        def dfs():
            s = [(root, -1, 0)]
            while s:
                v, par, depth = s.pop()
                parents[0][v] = par
                dists[v] = depth
                for to in adj[v]:
                    if to == par: continue
                    s.append((to, v, depth+1))

        dfs()
        for k in range(sz-1):
            for v in range(n):
                if parents[k][v] < 0: continue
                parents[k+1][v] = parents[k][parents[k][v]]

        self.parents = parents
        self.dists = dists

    def query(self, u: int, v: int) -> int:
        """二頂点 u, v の LCA"""
        if self.dists[u] < self.dists[v]:
            u, v = v, u

        sz = len(self.parents)
        # LCA までの距離を同じにする
        for k in range(sz):
            if (self.dists[u] - self.dists[v]) >> k & 1:
                u = self.parents[k][u]

        if u == v: return u

        assert self.dists[u] == self.dists[v]
        for k in reversed(range(sz)):
            if self.parents[k][u] != self.parents[k][v]:
                u = self.parents[k][u]
                v = self.parents[k][v]

        return self.parents[0][u]

    def distance(self, u: int, v: int) -> int:
        """二頂点 u, v の距離"""
        return self.dists[u] + self.dists[v] - 2 * self.dists[self.query(u, v)]

    def is_on_path(self, u: int, v: int, a: int) -> bool:
        """二頂点 u, v 上に頂点 a があるか"""
        return self.distance(u, a) + self.distance(a, v) == self.distance(u, v)

    def get_path(self, u: int, v: int) -> list[int]:
        """二頂点 u, v 間のパスを求める"""
        def up(k: int, par: int) -> list[int]:
            n = len(self.dists)
            path = [k]
            while path[-1] != par:
                assert len(path) <= n
                k = self.parents[0][k]
                path.append(k)

            return path

        p = self.query(u, v)
        a = up(u, p)  # u -> p
        b = up(v, p)  # v -> p

        return a[:-1] + b[::-1]


class EulerTourEdge:
    # n : 頂点数
    # adj : 隣接頂点 { 頂点: [(隣接頂点, 重み, インデックス)...] }
    def __init__(self, n, adj, root=0):
        vs = []  # 頂点 (訪問順)
        es = []  # 辺の重み (訪問順。葉への方向は正。根への方向は負)
        v2i = [INF] * n  # v2i[v] : 頂点 v が vs に現れる最初のインデックス
        e_in = [0] * (n-1)  # 辺 i の子方向への ws のインデックス
        e_out = [0] * (n-1)  # 辺 i の親方向への ws のインデックス

        def dfs(v, par):
            v2i[v] = len(vs)
            vs.append(v)
            for to, w, ind in adj[v]:
                if to == par: continue
                # 子への遷移
                e_in[ind] = len(es)
                es.append(w)
                dfs(to, v)
                # 親への遷移
                vs.append(v)
                e_out[ind] = len(es)
                es.append(-w)

        dfs(root, -1)

        def heads():
            return {k: [v[0] for v in v] for k, v in adj.items()}

        self.lca = LCA(n, heads(), root)
        self.segt = SegTree(lambda a, b: a+b, 0, es)
        self.v2i = v2i
        self.e_in = e_in
        self.e_out = e_out

    def change_edge_weight(self, e, w):
        """辺 e の重みを w に変更"""
        self.segt.set(self.e_in[e], w)
        self.segt.set(self.e_out[e], -w)

    def distance(self, u, v) -> int:
        """頂点 u, v の距離"""
        du = self.segt.prod(0, self.v2i[u])
        dv = self.segt.prod(0, self.v2i[v])
        da = self.segt.prod(0, self.v2i[self.lca.query(u, v)])
        return du + dv - 2 * da


from collections import defaultdict

INF = 1 << 62

T = int(input())
N, M = map(int, input().split())

edges = []
for _ in range(M):
    U, V, W = map(int, input().split())
    U -= 1
    V -= 1
    edges.append((U, V, W))


def solve_undirected():
    sorted_edges = sorted(edges, key=lambda x: x[2])
    uf = UnionFind(N)
    adj = defaultdict(list)
    used = set()

    for i, (u, v, w) in enumerate(sorted_edges):
        if not uf.is_same(u, v):
            uf.unite(u, v)
            adj[u].append((v, w, i))
            adj[v].append((u, w, i))
            used.add(i)

    t = EulerTourEdge(N, adj)
    res = INF
    for i, (u, v, w) in enumerate(sorted_edges):
        if i in used: continue

        d = t.distance(u, v) + w
        res = min(res, d)

    if res == INF:
        return -1
    return res


def solve_directed():
    return 0


if T == 0:
    ans = solve_undirected()
    print(ans)
else:
    ans = solve_directed()
    print(ans)
0