感觉写的没啥问题,但一直 50 pts,dalao 们能不能帮忙调一下 qwq
#include<vector>
#include<cstdio>
#include<iostream>
#include<algorithm>
using namespace std;
inline int read()
{
int x = 0, f = 1; char ch = getchar();
while (ch < '0' || ch > '9') { if (ch == '-') f = -1; ch = getchar(); }
while (ch >= '0' && ch <= '9') { x = x * 10 + ch - 48; ch = getchar(); }
return x * f;
}
const int MAXN = 3e5 + 10;
struct edge
{
int nxt, to;
}e[MAXN << 1];
int head[MAXN], cnt;
void add_edge(int u, int v)
{
e[++cnt].nxt = head[u];
head[u] = cnt;
e[cnt].to = v;
}
struct node
{
int s, t, lca;
}P[MAXN];
vector<int>v1[MAXN];
vector<int>v2[MAXN];
int n, m, w[MAXN], dep[MAXN], anc[MAXN][30];
int buc1[MAXN], buc2[MAXN * 2], ans[MAXN];
void dfs1(int x, int fa, int d)
{
dep[x] = d;
for (int i = 1; i <= 20; i++)
anc[x][i] = anc[anc[x][i - 1]][i - 1];
for (int i = head[x]; i; i = e[i].nxt)
{
int v = e[i].to;
if (v == fa) continue;
anc[v][0] = x;
dfs1(v, x, d + 1);
}
}
void dfs2(int x, int fa)
{
int res1 = buc1[w[x] + dep[x]];
int res2 = buc2[w[x] - dep[x] + MAXN];
int det = 0;
if (!v1[x].empty())
{
buc1[dep[x]] += v1[x].size();
// printf("*%d\n", buc[1]);
}
if (!v2[x].empty())
for (int i = 0; i < v2[x].size(); i++)
{
int pos = v2[x][i];
buc2[dep[P[pos].s] - 2 * dep[P[pos].lca] + MAXN]++;
}
for (int i = head[x]; i; i = e[i].nxt)
{
int v = e[i].to;
if (v == fa) continue;
dfs2(v, x);
}
int res3 = buc1[w[x] + dep[x]];
int res4 = buc2[w[x] - dep[x] + MAXN];
ans[x] += res3 + res4 - res1 - res2;
}
int LCA(int x, int y)
{
if (dep[x] > dep[y]) swap(x, y);
int len = dep[y] - dep[x];
for (int i = 20; i >= 0; i--)
if ((1 << i) & len) y = anc[y][i];
if (x == y) return x;
for (int i = 20; i >= 0; i--)
if (anc[x][i] != anc[y][i])
x = anc[x][i], y = anc[y][i];
return anc[x][0];
}
int main()
{
n = read(); m = read();
for (int i = 1; i < n; i++)
{
int u = read(), v = read();
add_edge(u, v);
add_edge(v, u);
}
for (int i = 1; i <= n; i++)
w[i] = read();
dfs1(1, 0, 0);
for (int i = 1; i <= m; i++)
{
P[i].s = read(); P[i].t = read();
P[i].lca = LCA(P[i].s, P[i].t);
v1[P[i].s].push_back(i);
v2[P[i].t].push_back(i);
if (dep[P[i].lca] + w[P[i].lca] == dep[P[i].s]) ans[P[i].lca]--;
}
dfs2(1, 0);
for (int i = 1; i <= n; i++)
printf("%d ", ans[i]);
return 0;
}