結果
| 問題 | No.3608 Golden Steiner Tree |
| コンテスト | |
| ユーザー |
👑 |
| 提出日時 | 2026-07-31 01:30:51 |
| 言語 | PyPy3 (7.3.17) |
| 結果 |
WA
|
| 実行時間 | - |
| コード長 | 12,553 bytes |
| 記録 | |
| コンパイル時間 | 265 ms |
| コンパイル使用メモリ | 96,108 KB |
| 実行使用メモリ | 292,452 KB |
| 最終ジャッジ日時 | 2026-07-31 20:53:59 |
| 合計ジャッジ時間 | 18,748 ms |
|
ジャッジサーバーID (参考情報) |
judge1_0 / judge3_0 |
(要ログイン)
| ファイルパターン | 結果 |
|---|---|
| sample | WA * 1 |
| other | AC * 1 WA * 11 TLE * 3 -- * 5 |
ソースコード
# https://github.com/tatyam-prime/SortedSet/blob/main/SortedSet.py
import math
from bisect import bisect_left, bisect_right
from collections.abc import Iterable, Iterator
from typing import Generic, TypeVar
class EulerTour:
n: int
tree: list[list[int]]
_result: list[int] | None = None
def __init__(self, n: int):
self.n = n
self.tree = [[] for _ in range(n)]
def add_edge(self, u: int, v: int):
self.tree[u].append(v)
self.tree[v].append(u)
def build(self, root: int = 0):
self._dfs(root)
def _dfs(self, root: int):
if self._result:
return
result: list[int] = []
IN, OUT = 0, 1
stack = [(root, float("NaN"), IN)]
while stack:
curr, prev, order = stack.pop()
if order == IN:
result.append(curr)
stack.append((curr, prev, OUT))
for to in self.tree[curr]:
if to == prev:
continue
stack.append((to, curr, IN))
else:
result.append(curr)
self._result = result
def order(self) -> list[int]:
assert self._result is not None
return self._result
def begin_end(self) -> tuple[list[int], list[int]]:
assert self._result is not None
begin = [-1] * self.n
end = [-1] * self.n
for i, v in enumerate(self._result):
if begin[v] == -1:
begin[v] = i
else:
end[v] = i
return begin, end
class LowestCommonAncestor:
"""
Lowest Common Ancestor (LCA)
---
木に対する最小共通祖先を求めるデータ構造
"""
def __init__(self, n: int):
"""\
木の頂点数 n を指定して初期化する
Parameters:
n (int): 木の頂点数
"""
self._n = n
self._logn = n.bit_length() + 1
self._depth = [0] * self._n
self._distance = [0] * self._n
self._ancestor = [-1 for _ in range(self._n * self._logn)]
self._edges = [[] for _ in range(self._n)]
def add_edge(self, u: int, v: int, w: int = 1):
"""\
u, v 間に重み w の辺を追加する
Parameters:
u (int): 辺の片方の頂点
v (int): 辺のもう片方の頂点
w (int): 辺の重み
"""
self._edges[u].append((v, w))
self._edges[v].append((u, w))
def build(self, root: int = 0):
"""\
根を root にした木を構築する
Parameters:
root (int): 根の頂点番号
"""
stack = [root]
while stack:
now = stack.pop()
for to, w in self._edges[now]:
if self._ancestor[to] == now or self._ancestor[now] == to:
continue
self._ancestor[to] = now
self._depth[to] = self._depth[now] + 1
self._distance[to] = self._distance[now] + w
stack.append(to)
for k in range(1, self._logn):
for i in range(self._n):
if self._ancestor[(k - 1) * self._n + i] == -1:
self._ancestor[k * self._n + i] = -1
else:
double = (k - 1) * self._n + self._ancestor[(k - 1) * self._n + i]
self._ancestor[k * self._n + i] = self._ancestor[double]
def lca(self, u: int, v: int) -> int:
"""\
u, v の最小共通祖先を求める
Parameters:
u (int): 頂点 u
v (int): 頂点 v
Returns:
lca (int): u, v の最小共通祖先
"""
# u の深さを v の深さ以下になるよう調整する
if self._depth[u] > self._depth[v]:
u, v = v, u
# v の深さを u に合わせる
for k in range(self._logn - 1, -1, -1):
if ((self._depth[v] - self._depth[u]) >> k) & 1 == 1:
v = self._ancestor[k * self._n + v]
# この時点で一致すれば、それが解
if u == v:
return u
# u, v がギリギリ一致しないよう親方向に辿る
for k in range(self._logn - 1, -1, -1):
if self._ancestor[k * self._n + u] != self._ancestor[k * self._n + v]:
u = self._ancestor[k * self._n + u]
v = self._ancestor[k * self._n + v]
# 最後に 1 ステップ親方向に辿った頂点が解
return self._ancestor[u]
# u, v (0-indexed) の距離を求める
def distance(self, u: int, v: int) -> int:
"""\
u, v 間の距離を求める
Parameters:
u (int): 頂点 u
v (int): 頂点 v
Returns:
dist (int): u, v 間の最短距離の長さ
"""
return self._distance[u] + self._distance[v] - 2 * self._distance[self.lca(u, v)]
# v の親を求める
def parent(self, v: int) -> int:
"""\
v の親を求める
Parameters:
v (int): 頂点 v
Returns:
parent (int): 頂点 v の親
"""
return self._ancestor[v]
def ancestor(self, v: int, gen: int) -> int:
if self._depth[v] < gen:
return None
curr = v
for i in range(self._logn):
if (gen >> i) & 1 == 1:
curr = self._ancestor[i * self._n + curr]
return curr
T = TypeVar("T")
class SortedSet(Generic[T]):
BUCKET_RATIO = 16
SPLIT_RATIO = 24
def __init__(self, a: Iterable[T] = []) -> None:
"Make a new SortedSet from iterable. / O(N) if sorted and unique / O(N log N)"
a = list(a)
n = len(a)
if any(a[i] > a[i + 1] for i in range(n - 1)):
a.sort()
if any(a[i] >= a[i + 1] for i in range(n - 1)):
a, b = [], a
for x in b:
if not a or a[-1] != x:
a.append(x)
n = self.size = len(a)
num_bucket = int(math.ceil(math.sqrt(n / self.BUCKET_RATIO)))
self.a = [a[n * i // num_bucket : n * (i + 1) // num_bucket] for i in range(num_bucket)]
def __iter__(self) -> Iterator[T]:
for i in self.a:
for j in i:
yield j
def __reversed__(self) -> Iterator[T]:
for i in reversed(self.a):
for j in reversed(i):
yield j
def __eq__(self, other) -> bool:
return list(self) == list(other)
def __len__(self) -> int:
return self.size
def __repr__(self) -> str:
return "SortedSet" + str(self.a)
def __str__(self) -> str:
s = str(list(self))
return "{" + s[1 : len(s) - 1] + "}"
def _position(self, x: T) -> tuple[list[T], int, int]:
"return the bucket, index of the bucket and position in which x should be. self must not be empty."
for i, a in enumerate(self.a):
if x <= a[-1]:
break
return (a, i, bisect_left(a, x))
def __contains__(self, x: T) -> bool:
if self.size == 0:
return False
a, _, i = self._position(x)
return i != len(a) and a[i] == x
def add(self, x: T) -> bool:
"Add an element and return True if added. / O(√N)"
if self.size == 0:
self.a = [[x]]
self.size = 1
return True
a, b, i = self._position(x)
if i != len(a) and a[i] == x:
return False
a.insert(i, x)
self.size += 1
if len(a) > len(self.a) * self.SPLIT_RATIO:
mid = len(a) >> 1
self.a[b : b + 1] = [a[:mid], a[mid:]]
return True
def _pop(self, a: list[T], b: int, i: int) -> T:
ans = a.pop(i)
self.size -= 1
if not a:
del self.a[b]
return ans
def discard(self, x: T) -> bool:
"Remove an element and return True if removed. / O(√N)"
if self.size == 0:
return False
a, b, i = self._position(x)
if i == len(a) or a[i] != x:
return False
self._pop(a, b, i)
return True
def lt(self, x: T) -> T | None:
"Find the largest element < x, or None if it doesn't exist."
for a in reversed(self.a):
if a[0] < x:
return a[bisect_left(a, x) - 1]
def le(self, x: T) -> T | None:
"Find the largest element <= x, or None if it doesn't exist."
for a in reversed(self.a):
if a[0] <= x:
return a[bisect_right(a, x) - 1]
def gt(self, x: T) -> T | None:
"Find the smallest element > x, or None if it doesn't exist."
for a in self.a:
if a[-1] > x:
return a[bisect_right(a, x)]
def ge(self, x: T) -> T | None:
"Find the smallest element >= x, or None if it doesn't exist."
for a in self.a:
if a[-1] >= x:
return a[bisect_left(a, x)]
def __getitem__(self, i: int) -> T:
"Return the i-th element."
if i < 0:
for a in reversed(self.a):
i += len(a)
if i >= 0:
return a[i]
else:
for a in self.a:
if i < len(a):
return a[i]
i -= len(a)
raise IndexError
def pop(self, i: int = -1) -> T:
"Pop and return the i-th element."
if i < 0:
for b, a in enumerate(reversed(self.a)):
i += len(a)
if i >= 0:
return self._pop(a, ~b, i)
else:
for b, a in enumerate(self.a):
if i < len(a):
return self._pop(a, b, i)
i -= len(a)
raise IndexError
def index(self, x: T) -> int:
"Count the number of elements < x."
ans = 0
for a in self.a:
if a[-1] >= x:
return ans + bisect_left(a, x)
ans += len(a)
return ans
def index_right(self, x: T) -> int:
"Count the number of elements <= x."
ans = 0
for a in self.a:
if a[-1] > x:
return ans + bisect_right(a, x)
ans += len(a)
return ans
LIMIT = 10**12
def golden(i):
ok, ng = LIMIT, -1
while abs(ok - ng) > 1:
mid = (ok + ng) // 2
if i**2 * 5 <= (2 * mid - i) ** 2:
ok = mid
else:
ng = mid
return ok
def golden_2(i):
ok, ng = LIMIT, -1
while abs(ok - ng) > 1:
mid = (ok + ng) // 2
if i**2 * 5 <= (2 * mid - 3 * i) ** 2:
ok = mid
else:
ng = mid
return ok
N, R, B = [int(s) for s in input().split()]
red_edges = []
blue_edges = []
for u in range(1, N + 1):
v = golden(u)
if u < v <= N:
red_edges.append((u - 1, v - 1))
v = golden_2(u)
if u < v <= N:
blue_edges.append((u - 1, v - 1))
euler = EulerTour(N)
for u, v in red_edges + blue_edges:
euler.add_edge(u, v)
euler.build()
begin, end = euler.begin_end()
inv_begin = {begin[i]: i for i in range(N)}
sset = SortedSet()
red = LowestCommonAncestor(N)
blue = LowestCommonAncestor(N)
for u, v in red_edges:
red.add_edge(u, v, 1)
blue.add_edge(u, v, 0)
for u, v in blue_edges:
red.add_edge(u, v, 0)
blue.add_edge(u, v, 1)
red.build()
blue.build()
ans_red = 0
ans_blue = 0
Q = int(input())
for _ in range(Q):
cmd, x = [int(s) for s in input().split()]
if cmd == 1:
x -= 1
if len(sset) == 0:
print(0)
sset.add(begin[x])
continue
l = sset.lt(begin[x])
if l is None:
l = sset[-1]
r = sset.gt(begin[x])
if r is None:
r = sset[0]
l, r = inv_begin[l], inv_begin[r]
if begin[x] in sset:
ans_red -= red.distance(l, x) + red.distance(x, r) - red.distance(l, r)
ans_blue -= blue.distance(l, x) + blue.distance(x, r) - blue.distance(l, r)
sset.discard(begin[x])
else:
ans_red += red.distance(l, x) + red.distance(x, r) - red.distance(l, r)
ans_blue += blue.distance(l, x) + blue.distance(x, r) - blue.distance(l, r)
sset.add(begin[x])
elif cmd == 2:
R = x
else:
B = x
print((ans_red * R + ans_blue * B) // 2)