小猪佩奇爬树

GitHub跳转原题关系图返回列表

把同色节点的路径约束化为 Steiner 子树直径,并用两侧组件计数路径数量。

OJ: luogu

题目 ID: P5588

难度:提高

标签:树的直径LCA树上计数python

日期: 2026-07-17 02:00

题意

对每种颜色,统计路径恰好经过该颜色全部节点的点对数量。

思路

若同色节点不共线,则没有路径能全部经过,答案为 0。否则同色节点都在它们直径端点 x,y 的路径上,路径经过全部节点等价于经过 x-y。切断 xy 方向的第一条边,两个外侧组件大小相乘就是答案;没有该颜色时所有点对都合法。

Python 知识

  • 颜色节点用“同色链表”存储,避免为百万种颜色创建百万个 Python 列表。
  • array 保存边、父亲、祖先和答案,控制百万节点内存。
  • 分块输出一百万行结果,避免构造巨大字符串列表。

代码

python
import sys
from array import array


input = sys.stdin.buffer.readline
n = int(input())
colors = array("i", [0])
colors.extend(map(int, input().split()))
head = array("i", [-1]) * (n + 1)
to = array("i", [0]) * (2 * n - 2)
next_edge = array("i", [0]) * (2 * n - 2)
edge_count = 0


def add_edge(u, v):
    global edge_count
    to[edge_count] = v
    next_edge[edge_count] = head[u]
    head[u] = edge_count
    edge_count += 1


for _ in range(n - 1):
    u, v = map(int, input().split())
    add_edge(u, v)
    add_edge(v, u)

parent = array("i", [0]) * (n + 1)
depth = array("i", [0]) * (n + 1)
order = array("i", [1])
index = 0
while index < n:
    node = order[index]
    index += 1
    edge = head[node]
    while edge != -1:
        neighbor = to[edge]
        if neighbor != parent[node]:
            parent[neighbor] = node
            depth[neighbor] = depth[node] + 1
            order.append(neighbor)
        edge = next_edge[edge]
subtree = array("i", [1]) * (n + 1)
for node in reversed(order[1:]):
    subtree[parent[node]] += subtree[node]

ancestors = [parent]
for _ in range(1, n.bit_length()):
    previous = ancestors[-1]
    ancestors.append(array("i", (previous[previous[node]] for node in range(n + 1))))


def jump(node, steps):
    bit = 0
    while steps:
        if steps & 1:
            node = ancestors[bit][node]
        steps >>= 1
        bit += 1
    return node


def lca(x, y):
    if depth[x] < depth[y]:
        x, y = y, x
    x = jump(x, depth[x] - depth[y])
    if x == y:
        return x
    for level in range(len(ancestors) - 1, -1, -1):
        if ancestors[level][x] != ancestors[level][y]:
            x = ancestors[level][x]
            y = ancestors[level][y]
    return parent[x]


def distance(x, y):
    ancestor = lca(x, y)
    return depth[x] + depth[y] - 2 * depth[ancestor]


color_head = array("i", [0]) * (n + 1)
next_same = array("i", [0]) * (n + 1)
color_count = array("i", [0]) * (n + 1)
present_colors = array("i")
for node in range(1, n + 1):
    color = colors[node]
    if color_count[color] == 0:
        present_colors.append(color)
    next_same[node] = color_head[color]
    color_head[color] = node
    color_count[color] += 1
all_pairs = n * (n - 1) // 2
answer = array("q", [all_pairs]) * (n + 1)


def side_size(x, y):
    """Size of the component containing x after cutting the first edge x -> y."""
    ancestor = lca(x, y)
    if ancestor == x:
        child = jump(y, depth[y] - depth[x] - 1)
        return n - subtree[child]
    return subtree[x]


for color in present_colors:
    if color_count[color] == 1:
        x = color_head[color]
        excluded = (n - subtree[x]) * (n - subtree[x] - 1) // 2
        edge = head[x]
        while edge != -1:
            child = to[edge]
            if parent[child] == x:
                size = subtree[child]
                excluded += size * (size - 1) // 2
            edge = next_edge[edge]
        answer[color] = all_pairs - excluded
        continue
    first = color_head[color]
    second = first
    best_distance = -1
    node = first
    while node:
        current_distance = distance(first, node)
        if current_distance > best_distance:
            best_distance = current_distance
            second = node
        node = next_same[node]
    third = second
    best_distance = -1
    node = first
    while node:
        current_distance = distance(second, node)
        if current_distance > best_distance:
            best_distance = current_distance
            third = node
        node = next_same[node]
    diameter_length = distance(second, third)
    node = first
    lies_on_path = True
    while node:
        if distance(second, node) + distance(node, third) != diameter_length:
            lies_on_path = False
            break
        node = next_same[node]
    if lies_on_path:
        answer[color] = side_size(second, third) * side_size(third, second)
    else:
        answer[color] = 0

output = sys.stdout.write
buffer = []
for color in range(1, n + 1):
    buffer.append(str(answer[color]))
    if len(buffer) == 8192:
        output("\n".join(buffer) + "\n")
        buffer.clear()
output("\n".join(buffer))

原有 C++ 版本仍保留:

cpp
/**
 * Author by Rainboy blog: https://rainboylv.com github: https://github.com/rainboylvx
 * rbook: -> https://rbook.roj.ac.cn  https://rbook2.roj.ac.cn
 * rainboy的学习导航网站: https://idx.roj.ac.cn
 * create_at: 2026-07-17 01:04
 * update_at: 2026-07-17 01:04
 */
#include <bits/stdc++.h>
using namespace std;

int main() {
    ios::sync_with_stdio(false);
    cin.tie(nullptr);

    return 0;
}

复杂度

总节点参与两次同色最远点扫描,时间约 O(n log n),空间 O(n log n)

总结

“经过一组点”先看这些点是否共线;共线后只需关注两个端点和两侧组件。