WA on #6 #7
#include<bits/stdc++.h>
using namespace std;
int n,m;
int tmp[300005],ans[300005];
struct seg{
#define lson lc[x]
#define rson rc[x]
#define mid ((l+r)>>1)
int a[300005*40],lc[300005*40],rc[300005*40],tot;
void addnode(int &x){
if(x==0)x=++tot;
}
void upd(int p,int op,int &x,int l,int r){
addnode(x);
if(l==r){
a[x]+=op;
return;
}
if(p<=mid)upd(p,op,lson,l,mid);
else upd(p,op,rson,mid+1,r);
}
int get(int p,int &x,int l,int r){
if(l==r)return a[x];
if(p<=mid)return get(p,lson,l,mid);
else return get(p,rson,mid+1,r);
}
void mer(int &x,int &y,int l,int r){
if(x==0){
x=y;
return;
}
if(y==0){
return;
}
a[x]+=a[y];
if(l==r)return;
mer(lc[x],lc[y],l,mid);
mer(rc[x],rc[y],mid+1,r);
}
}hitoi,bocchi;
int root[300005],rt[300005];
vector<int>mp[300005];
struct LCA{
int f[300005][23],dep[300005];
void init(int x,int fa){
f[x][0]=fa;
dep[x]=dep[fa]+1;
for(int i=1;i<=20;i++){
f[x][i]=f[f[x][i-1]][i-1];
}
for(int i:mp[x]){
if(i==fa)continue;
init(i,x);
}
}
int lca(int a,int b){
if(dep[a]<dep[b])swap(a,b);
for(int i=20;i>=0;i--){
if(dep[f[a][i]]>=dep[b])a=f[a][i];
}
if(a==b)return a;
for(int i=20;i>=0;i--){
if(f[a][i]!=f[b][i])a=f[a][i],b=f[b][i];
}
return f[a][0];
}
}cul;
struct node_to_lca{
void add(int p,int op,int x){
hitoi.upd(p,op,root[x],1,n);
}
void dfs1(int x,int f){
for(int i:mp[x]){
if(i==f)continue;
dfs1(i,x);
hitoi.mer(root[x],root[i],1,n);
}
ans[x]+=hitoi.get(cul.dep[x]+tmp[x],root[x],1,n);
}
}oper1;
struct lca_to_node{
void add(int p,int op,int x){
bocchi.upd(p+n,op,rt[x],1,n+n);
}
void dfs2(int x,int f){
for(int i:mp[x]){
if(i==f)continue;
dfs2(i,x);
bocchi.mer(rt[x],rt[i],1,n+n);
}
ans[x]+=bocchi.get(cul.dep[x]-tmp[x]+n,rt[x],1,n+n);
}
}oper2;
int main(){
freopen("P1600.in","r",stdin);
freopen("out.txt","w",stdout);
scanf("%d%d",&n,&m);
for(int i=1;i<n;i++){
int u,v;
scanf("%d%d",&u,&v);
mp[u].push_back(v);
mp[v].push_back(u);
}
cul.init(1,0);
for(int i=1;i<=n;i++)scanf("%d",tmp+i);
for(int i=1;i<=m;i++){
int s,t;
scanf("%d%d",&s,&t);
int ui=cul.lca(s,t);
if(cul.dep[s]-cul.dep[ui]==tmp[ui])ans[ui]++;
oper1.add(cul.dep[s],1,s);
oper1.add(cul.dep[s],-1,ui);
oper2.add(cul.dep[ui]-(cul.dep[s]-cul.dep[ui]),1,t);
oper2.add(cul.dep[ui]-(cul.dep[s]-cul.dep[ui]),-1,ui);
}
oper1.dfs1(1,0);
// for(int i=1;i<=n;i++)printf("%d ",ans[i]);
// cout<<endl;
oper2.dfs2(1,0);
for(int i=1;i<=n;i++)printf("%d ",ans[i]);
return 0;
}
采用的是线段树合并的写法。将贡献拆成 s→lca lca→t 再特判 lca