CF1794E Labeling the Tree with Distances

题目大意:给定一棵 \(n\) 个节点的树和一个长度为 \(n-1\) 的序列,对于一个树上的点 \(u\),你需要把序列中的数填到每一个节点上,剩下一个可以随便填,使得每个点上的数是它到 \(u\) 的距离。请求出所有的 \(u\)

由于 \(n \le 2\times 10^5\),我们对于每个点,最简单的办法就是去重新更新每一个点,复杂度 \(O(n^2)\)

所以每次都更新每个点是不行的,我们考虑转化为对每个点维护一个值,就考虑 \(D_u\) 把其他的所有点到 \(u\) 的距离加起来的值。但是这样的话很容易冲突,所以我们考虑哈希,使得 \(D_u=\sum_{v∈S} base^{d_v}\),这样冲突概率就大大降低了,在判断能否用序列表示的时候只需要看 \(D_u-S\) 剩下的是否是一个数即可。(\(S\) 是序列通过哈希处理后的总数)。

然后看看怎么换根。对于一个 \(v\),假设它的父节点是 \(u\),对于 \(v\) 子树外的节点,他们的深度都 \(+1\),而对于 \(v\) 子树内就都 \(-1\)。这样维护一下 \(v\) 子树外的值就好了。

#include <bits/stdc++.h>
#define ll long long
using namespace std;
const int N = 2e5 + 10;
const ll mod1 = 1e9 + 7, bas1 = 13331;
const ll mod2 = 998244353, bas2 = 1145141;
ll pw1[N], pw2[N];
struct node {ll s; int id;} pw[N];
bool cmp(node a, node b) {return a.s < b.s;}

vector<int> g[N];
ll inD1[N], inD2[N], D1[N], D2[N];
void dfs1(int u, int fa) {
	inD1[u] = inD2[u] = 1;
	for(auto v : g[u]) if(v != fa) {
		dfs1(v, u);
		inD1[u] = (inD1[u] + inD1[v] * bas1) % mod1;
		inD2[u] = (inD2[u] + inD2[v] * bas2) % mod2;
	}
}
void dfs2(int u, int fa) {
	for(auto v : g[u]) if(v != fa) {
		D1[v] = ((D1[u] - inD1[v] * bas1 % mod1 + mod1) * bas1 + inD1[v]) % mod1;
		D2[v] = ((D2[u] - inD2[v] * bas2 % mod2 + mod2) * bas2 + inD2[v]) % mod2;
		dfs2(v, u);
	}
}
vector<int> ans;
void solve() {
	int n; scanf("%d", &n);
	pw1[0] = pw2[0] = 1;
	pw[0] = {1, 0};
	for(int i = 1; i < n; i++) {
		pw1[i] = pw1[i - 1] * bas1 % mod1;
		pw2[i] = pw2[i - 1] * bas2 % mod2;
		pw[i] = {pw1[i], i};
	}
	sort(pw, pw + n, cmp);
	ll H1 = 0, H2 = 0;
	for(int i = 1; i < n; i++) {
		int x; scanf("%d", &x);
		H1 = (H1 + pw1[x]) % mod1;
		H2 = (H2 + pw2[x]) % mod2;
	}
	for(int i = 1; i < n; i++) {
		int u, v; scanf("%d%d", &u, &v);
		g[u].push_back(v), g[v].push_back(u);
	}
	dfs1(1, 0); 
	D1[1] = inD1[1]; D2[1] = inD2[1];
	dfs2(1, 0);
	for(int u = 1; u <= n; u++) {
		ll H = (D1[u] - H1 + mod1) % mod1;
		int l = 0, r = n - 1, pos = n;
		while(l <= r) {
			int mid = l + r >> 1;
			if(pw[mid].s >= H) r = mid - 1, pos = mid;
			else l = mid + 1;
		}
		while(pos < n && pw[pos].s == H) {
			if(pw2[pw[pos].id] == (D2[u] - H2 + mod2) % mod2) {
				ans.push_back(u); break;
			}
			++pos;
		}
	}
	printf("%d\n", ans.size());
	for(int x : ans) printf("%d ", x);
}
int main() {
    int T; T = 1;
    while(T--) solve();
    return 0;
}
posted @ 2026-08-08 10:07  OIerYang  阅读(3)  评论(0)    收藏  举报