警示后人
查看原帖
警示后人
658786
STUDENT00楼主2022/11/19 13:24

计算过程中可能出现负数,所以在计算最大值时,要将初始值设为 109-10^9

详见代码:

#include<bits/stdc++.h>
#define N 30010
using namespace std;
int n,q,s[N],d[N],f[N],size[N],son[N],top[N],id[N],rk[N],cnt,sum[N<<2],maxs[N<<2];
vector<int> w[N],v[N]; 
bool vis[N];
void dfs1(int now){
	vis[now]=size[now]=1;
	d[now]=d[f[now]]+1;
	for(int i=0;i<w[now].size();i++){
		int t=w[now][i];
		if(!vis[t]){
			v[now].push_back(t);
			vis[t]=1;
			f[t]=now;
			dfs1(t);
			size[now]+=size[t];
			if(size[t]>size[son[now]]) son[now]=t;
		}
	}
}
void dfs2(int now,int t){
	top[now]=t;
	id[now]=++cnt;
	rk[cnt]=now;
	if(son[now]) dfs2(son[now],t);
	for(int i=0;i<v[now].size();i++){
		int p=v[now][i];
		if(p!=son[now]) dfs2(p,p);
	}
}
void push_up(int rt){
	sum[rt]=sum[rt<<1]+sum[rt<<1|1];
	maxs[rt]=max(maxs[rt<<1],maxs[rt<<1|1]);
}
void build(int l,int r,int rt){
	if(l==r){
		sum[rt]=maxs[rt]=s[rk[l]];
		return;
	}
	int mid=l+r>>1;
	build(l,mid,rt<<1);
	build(mid+1,r,rt<<1|1);
	push_up(rt);
	return; 
}
void update(int l,int r,int rt,int a,int b){
	if(l==r){
		maxs[rt]=sum[rt]=b;
		return;
	}
	int mid=l+r>>1;
	if(a<=mid) update(l,mid,rt<<1,a,b);
	else update(mid+1,r,rt<<1|1,a,b);
	push_up(rt);
	return;
}
int query(int l,int r,int rt,int a,int b,int op){
	if(a<=l&&b>=r){
		if(op==1) return maxs[rt];
		else return sum[rt];
	}
	int mid=l+r>>1,ans;
	if(op==1) ans=-1e9;
	else ans=0;
	if(a<=mid){
		if(op==1) ans=max(ans,query(l,mid,rt<<1,a,b,1));
		else ans+=query(l,mid,rt<<1,a,b,2);
	}
	if(b>mid){
		if(op==1) ans=max(ans,query(mid+1,r,rt<<1|1,a,b,1));
		else ans+=query(mid+1,r,rt<<1|1,a,b,2);
	}
	return ans;
}
int querys(int x,int y,int op){
	int fx=top[x],fy=top[y],ans;
	if(op==1) ans=-1e9;
	else ans=0;
	while(fx!=fy){
		if(d[fx]<d[fy]){
			swap(x,y);
			swap(fx,fy);
		}
		if(op==1) ans=max(ans,query(1,n,1,id[fx],id[x],1));
		else ans+=query(1,n,1,id[fx],id[x],2);
		x=f[fx];
		fx=top[x];
	}
	if(id[x]>id[y]) swap(x,y);
	if(op==1) ans=max(ans,query(1,n,1,id[x],id[y],1));
	else ans+=query(1,n,1,id[x],id[y],2);
	return ans;
}
int main(){
	memset(maxs,128,sizeof(maxs));
	scanf("%d",&n);
	for(int i=1;i<n;i++){
		int a,b;
		scanf("%d%d",&a,&b);
		w[a].push_back(b);
		w[b].push_back(a);
	}
	for(int i=1;i<=n;i++) scanf("%d",&s[i]);
	dfs1(1);
	dfs2(1,1);
	build(1,n,1);
	scanf("%d",&q);
	while(q--){·
		char c[10];
		int x,y;
		scanf("%s%d%d",c,&x,&y);
		if(c[0]=='C') update(1,n,1,id[x],y);
		else{
			if(c[1]=='M') printf("%d\n",querys(x,y,1));
			else printf("%d\n",querys(x,y,2));
		}
	}
	return 0;
}
2022/11/19 13:24
加载中...