这边代码条理清晰,恳求大佬帮忙瞅一眼吧!
可以适当酬谢。
#include<bits/stdc++.h>
using namespace std;
const int maxn=4e5+10;
const int mod=1e9+7;
#define inf 1e9
inline int read(){
int x=0,f=1;char c=getchar();
while(c<'0'||c>'9'){if(c=='-')f=-1;c=getchar();}
while(c>='0'&&c<='9'){x=(x<<1)+(x<<3)+c-'0';c=getchar();}
return x*f;
}
#define ll long long
int n,m;ll a[maxn],siz[maxn],ans,Mx[maxn];
vector<int>G[maxn];
struct node{int ch[2],fa;}tr[maxn];
#define pb push_back
#define fa(x) tr[x].fa
#define lc(x) tr[x].ch[0]
#define rc(x) tr[x].ch[1]
inline bool isrt(int x){return lc(fa(x))!=x&&rc(fa(x))!=x;}
inline int chk(int x){return rc(fa(x))==x;}
inline void rotate(int x){
int y=fa(x),k=chk(x),t=tr[x].ch[!k],z=fa(y);
if(!isrt(y))tr[z].ch[chk(y)]=x;tr[x].ch[!k]=y;tr[y].ch[k]=t;
fa(x)=z;fa(y)=x;if(t)fa(t)=y;//pushup(y),pushup(x);
}
int st[maxn],top;ll Tr[maxn];
inline void splay(int x){
// int tmp=x;st[top=1]=tmp;
// while(!isrt(tmp))st[++top]=tmp=fa(tmp);
// while(top)pushdown(st[top]),top--;
for(;!isrt(x);rotate(x))
if(!isrt(fa(x)))rotate((chk(x)^chk(fa(x)))?x:fa(x));
//pushup(x);
}
//inline void access(int x){
// for(int y=0;x;y=x,x=fa(x))
// splay(x),rc(x)=y,pushup(x);
//}
inline int getson(int x){
splay(x);int tmp=rc(x);
if(!rc(x))return -1;
while(lc(tmp))tmp=lc(tmp);
return tmp;
}
inline int gettop(int x){
while(lc(x))x=lc(x);
return x;
}
int dfn[maxn],ti,sz[maxn];
inline void Add(int x,ll y){
for(;x<=n;x+=x&(-x))Tr[x]+=y;
}
inline ll Query(int x){
ll res=0;
for(;x;x-=x&(-x))res+=Tr[x];
return res;
}
inline ll getsiz(int x){
return Query(dfn[x]+sz[x]-1)-Query(dfn[x]-1);
}
inline void dfs(int x,int fa){
// printf("x=%d fa=%d\n",x,fa);
siz[x]=a[x];Mx[x]=a[x];fa(x)=fa;sz[x]=1;
dfn[x]=++ti;Add(dfn[x],a[x]);
for(auto t:G[x])if(t^fa){
dfs(t,x);siz[x]+=siz[t];sz[x]+=sz[t];
Mx[x]=max(Mx[x],siz[t]);
}ans+=siz[x]-max(1ll,2*Mx[x]-siz[x]);
for(auto t:G[x])if(t^fa)
if(2*siz[t]>siz[x])rc(x)=t;
}
int main(){
freopen("history.in","r",stdin);
freopen("history.out","w",stdout);
n=read(),m=read();
for(int i=1;i<=n;i++)a[i]=read();//puts("1");
for(int i=1,x,y;i<n;i++)
x=read(),y=read(),G[x].pb(y),G[y].pb(x);//puts("2");
dfs(1,0);printf("%lld\n",ans);
for(int i=1,x,y;i<=m;i++){
x=read(),y=read();ll v1=getsiz(x);
int t=getson(x);if(t!=-1){
ll v2=getsiz(t);ans-=2*(v1-v2);
if(2*v2>v1+y)ans+=2*(v1+y-v2);
else if(2*(a[x]+y)>v1+y)ans+=2*(v1-a[x]),rc(x)=0;
else ans+=v1+y-1,rc(x)=0;
}else if(2*a[x]<=v1){
ans-=v1-1;
if(2*(a[x]+y)>v1+y)ans+=2*(v1-a[x]);
else ans+=v1+y-1;
}
while(x){
splay(x);int fa=fa(x);
// printf("x=%d fa=%d ans=%lld\n",x,fa,ans);
if(!fa)break;int top=gettop(x);
ll v1=getsiz(fa)+y,v2=getsiz(top)+y;
int t=getson(fa);
// printf("top=%d v1=%lld v2=%lld t=%d\n",top,v1,v2,t);
if(2*v2>v1){
if(t!=-1)ans-=2*(v1-y-getsiz(t));
else if(2*a[fa]>v1-y)ans-=2*(v1-y-a[fa]);
else ans-=v1-y-1;
ans+=2*(v1-v2);rc(fa)=x;
}else{
if(t!=-1){
ll v3=getsiz(t);ans-=2*(v1-y-v3);
if(2*v3>v1)ans+=2*(v1-v3);
else if(2*a[fa]>v1)ans+=2*(v1-a[fa]);
else ans+=v1-1;
}else{
if(2*a[fa]>v1-y)ans-=2*(v1-y-a[fa]);
else ans-=v1-y-1;
if(2*a[fa]>v1)ans+=2*(v1-a[fa]);
else ans+=v1-1;
}
}x=fa;
}Add(dfn[x],y);a[x]+=y;
printf("%lld\n",ans);
}
return 0;
}