結果

問題 No.2337 Equidistant
コンテスト
ユーザー detteiuu
提出日時 2026-07-22 02:55:16
言語 PyPy3
(7.3.17)
コンパイル:
pypy3 -mpy_compile _filename_
実行:
pypy3 _filename_
結果
WA  
実行時間 -
コード長 3,895 bytes
記録
記録タグの例:
初AC ショートコード 純ショートコード 純主流ショートコード 最速実行時間
コンパイル時間 270 ms
コンパイル使用メモリ 96,364 KB
実行使用メモリ 245,312 KB
最終ジャッジ日時 2026-07-22 02:55:45
合計ジャッジ時間 27,516 ms
ジャッジサーバーID
(参考情報)
judge1_0 / judge2_0
このコードへのチャレンジ
(要ログイン)
ファイルパターン 結果
sample AC * 1
other AC * 2 WA * 26
権限があれば一括ダウンロードができます

ソースコード

diff #
raw source code

from sys import stdin
input = stdin.readline
from types import GeneratorType
from collections import deque

def bootstrap(f, stack=[]):
    def wrappedfunc(*args, **kwargs):
        if stack:
            return f(*args, **kwargs)
        to = f(*args, **kwargs)
        while True:
            if type(to) is GeneratorType:
                stack.append(to)
                to = next(to)
            else:
                stack.pop()
                if not stack:
                    break
                to = stack[-1].send(to)
        return to
    return wrappedfunc

class LCA:
    def __init__(self, G, root=0):
        V = len(G)
        self.bit_length = 1
        while 1<<self.bit_length < V:
            self.bit_length += 1
        self.parent = [[-1]*V for _ in range(self.bit_length)]
        self.depth = [-1]*V
        self.bfs(G, root)
        for i in range(self.bit_length-1):
            for j in range(V):
                if self.parent[i][j] != -1:
                    self.parent[i+1][j] = self.parent[i][self.parent[i][j]]
        
    def bfs(self, G, root):
        self.depth[root] = 0
        que = deque()
        que.append(root)
        while que:
            n = que.popleft()
            for v in G[n]:
                if self.depth[v] == -1:
                    self.depth[v] = self.depth[n]+1
                    self.parent[0][v] = n
                    que.append(v)
    
    def lca(self, a, b):
        if self.depth[a] < self.depth[b]:
            a, b = b, a
        for i in range(self.bit_length):
            if (self.depth[a]-self.depth[b]) & 1<<i:
                a = self.parent[i][a]
        if a == b:
            return a
        for i in reversed(range(self.bit_length)):
            if self.parent[i][a] != self.parent[i][b]:
                a = self.parent[i][a]
                b = self.parent[i][b]
        return self.parent[0][a]
    
    def dist(self, a, b):
        return self.depth[a]+self.depth[b]-self.depth[self.lca(a, b)]*2
    
    def is_ancestor(self, u, v):
        return self.lca(u, v) == u
    
    def kth_ancestor(self, n, k):
        for i in range(self.bit_length):
            if 1<<i & k:
                n = self.parent[i][n]
                if n == -1:
                    break
        return n
    
    def jump(self, u, v):
        if u == v:
            return u
        if self.lca(u, v) == u:
            return self.kth_ancestor(v, self.dist(u, v)-1)
        else:
            return self.parent[0][u]
    
    def jump_k(self, u, v, k):
        d = self.dist(u, v)
        if d <= k:
            return v
        lca = self.lca(u, v)
        if k <= self.depth[u]-self.depth[lca]:
            return self.kth_ancestor(u, k)
        else:
            return self.kth_ancestor(v, d-k)
    
    def on_path(self, u, v, x):
        return self.dist(u, x)+self.dist(x, v) == self.dist(u, v)

    def create_path(self, u, v, x):
        if self.on_path(u, v, x):
            return u, v
        if self.on_path(u, x, v):
            return u, x
        if self.on_path(v, x, u):
            return v, x
        return -1, -1

N, Q = map(int, input().split())
G = [[] for _ in range(N)]
for _ in range(N-1):
    u, v = map(int, input().split())
    u, v = u-1, v-1
    G[u].append(v)
    G[v].append(u)
query = [list(map(int, input().split())) for _ in range(Q)]

@bootstrap
def dfs(n, p):
    for v in G[n]:
        if v == p: continue
        yield dfs(v, n)
        size[n] += size[v]
    yield size[n]

size = [1]*N
dfs(0, -1)

lca = LCA(G)
for u, v in query:
    u, v = u-1, v-1
    lc = lca.lca(u, v)
    d = lca.dist(u, v)
    if d%2 == 1:
        print(0)
        continue
    m = lca.jump_k(u, v, d//2)
    if m == lc:
        print(N-size[lca.jump(m, u)]-size[lca.jump(m, v)])
    else:
        if lca.dist(u, m) < lca.dist(m, v):
            u, v = v, u
        print(N-size[lca.jump(m, v)]-(N-size[m]))
0