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;
}

浙公网安备 33010602011771号