結果

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

ソースコード

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 : 頂点数
    def __init__(self, n, edges, 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 のインデックス

        adj = defaultdict(list)
        for i, (u, v, w) in enumerate(edges):
            adj[u].append((v, w, i))
            adj[v].append((u, w, i))

        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: [to for to, _, _ 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
from heapq import heappush, heappop

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():
    res = INF

    sorted_edges = sorted(edges, key=lambda x: x[2])
    uf = UnionFind(N)
    adj = defaultdict(list)
    edge_used = set()
    for i, (u, v, w) in enumerate(sorted_edges):
        adj[u].append((v, w, i))
        adj[v].append((u, w, i))

        if not uf.is_same(u, v):
            uf.unite(u, v)
            edge_used.add(i)

    root_used = set()
    for i in range(N):  # 頂点 i を含む全域木
        root = uf.root(i)
        if root in root_used: continue
        root_used.add(root)

        i_edges = []
        for v in uf.group(i):
            for to, w, ind in adj[v]:
                if (v, to) in edge_used: continue
                if (to, v) in edge_used: continue
                if ind in edge_used:
                    i_edges.append((to, v, w))
                    edge_used.add((v, to))
                    edge_used.add((to, v))

        t = EulerTourEdge(N, i_edges, root=i)
        for i, (u, v, w) in enumerate(sorted_edges):
            if i in edge_used: continue
            if root != uf.root(u): continue
            if not uf.is_same(u, v): continue

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

    if res == INF:
        return -1
    return res


def solve_directed():
    res = INF

    adj = defaultdict(list)
    edge2w = {}
    for u, v, w in edges:
        adj[u].append((v, w))
        edge2w[u, v] = w

    for i in range(N):
        # 頂点 i を始点として、各頂点への最小経路を作る
        dists = [INF] * N
        dists[i] = 0
        q = [(0, i)]
        while q:
            d, v = heappop(q)
            if dists[v] != d: continue

            for to, w in adj[v]:
                nd = dists[v] + w
                if dists[to] > nd:
                    dists[to] = nd
                    heappush(q, (nd, to))

        # 各頂点から始点 i への有向辺があるなら閉路が存在する
        for j in range(N):
            if dists[j] != INF:
                if (j, i) in edge2w:
                    res = min(res, dists[j] + edge2w[j, i])

    if res == INF:
        return -1
    return res


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