萌新刚学树剖,此题20pts,AC #5 #6,求调qwq,悬赏关注
查看原帖
萌新刚学树剖,此题20pts,AC #5 #6,求调qwq,悬赏关注
607952
ZHANGGUIZHI楼主2023/1/8 10:18

改了很久还是没有看出问题,下面是代码

#include<bits/stdc++.h>
#define int long long
using namespace std;
int n,w[30002],q;
int head[30002],ver[60002],nex[60002],tot;//邻接表 
int dep[30002],fa[30002],size[30002],son[30002];
/*节点深度,节点的父亲,该节点及其子树的长度,该节点的重儿子*/
int top[30002],seg[30002],rev[30002],order;
/*该节点所在重路径的顶部节点,节点在线段树中的位置,对应的节点编号*/
string s;
void add(int x,int y){
	nex[++tot]=head[x];
	ver[tot]=y;
	head[x]=tot;
}
void dfs1(int u,int f){
	dep[u]=dep[f]+1;
	fa[u]=f;
	size[u]=1;
	for(int i=head[u];i;i=nex[i]){
		int v=ver[i];
		if(v==f)continue;
		dfs1(v,u);
		size[u]+=size[v];
		if(size[v]>size[son[u]])
		son[u]=v;
	}
}
void dfs2(int u,int t){
	top[u]=t;
	seg[u]=++order;
	rev[order]=u;
	if(son[u])dfs2(son[u],t);
	for(int i=head[u];i;i=nex[i]){
		int v=ver[i];
		if(v!=fa[u]&&v!=son[u])
		dfs2(v,v);
	}
}
struct node{
	int l,r,dat,sum;
}t[120002];
void build(int p,int l,int r){
	t[p].l=l,t[p].r=r;
	if(l==r){
		t[p].dat=w[rev[l]];
		t[p].sum=w[rev[l]];
		return ;
	}
	int mid=l+r>>1;
	build(p<<1,l,mid);
	build(p<<1|1,mid+1,r);
	t[p].sum=t[p<<1].sum+t[p<<1|1].sum;
	t[p].dat=max(t[p<<1].dat,t[p<<1|1].dat);
}
void change(int p,int x,int v){
	if(t[p].l==t[p].r&&t[p].l==x){
		t[p].dat=v;
		t[p].sum=v;
		return ;
	}
	int mid=t[p].l+t[p].r>>1;
	if(x<=mid)change(p<<1,x,v);
	else change(p<<1|1,x,v);
	t[p].dat=max(t[p<<1].dat,t[p<<1|1].dat);
	t[p].sum=t[p<<1].sum+t[p<<1|1].sum;
}
int askmax(int p,int l,int r){
	if(l>t[p].r||r<t[p].l)
	return 0;
	if(l<=t[p].l&&r>=t[p].r)
	return t[p].dat;
	int maxx=-2147483647;
	int mid=t[p].l+t[p].r>>1;
	if(l<=mid)maxx=max(maxx,askmax(p<<1,l,r));
	if(r>mid)maxx=max(maxx,askmax(p<<1|1,l,r));
	return maxx;
}
int askmax1(int x,int y){
	int maxn=-2147483647;
	while(top[x]!=top[y]){
		if(dep[top[x]]<dep[top[y]])
		swap(x,y);
		maxn=max(maxn,askmax(1,seg[top[x]],seg[x]));
		x=fa[top[x]];
	}
	if(dep[x]>dep[y])
	swap(x,y);
	maxn=max(maxn,askmax(1,seg[x],seg[y]));
	return maxn;
}
int asksum(int p,int l,int r){
	if(l>t[p].r||r<t[p].l)
	return 0;
	if(l<=t[p].l&&r>=t[p].r)
	return t[p].sum;
	int mid=t[p].l+t[p].r>>1,ans=0;
	if(l<=mid)ans=asksum(p<<1,l,r);
	if(r>mid)ans+=asksum(p<<1|1,l,r);
	return ans;
}
int asksum1(int x,int y){
	int res=0;
	while(top[x]!=top[y]){
		if(dep[top[x]]<dep[top[y]])
		swap(x,y);
		res+=asksum(1,seg[x],x);
		x=fa[top[x]];
	}
	if(dep[x]>dep[y])
	swap(x,y);
	res+=asksum(1,seg[x],seg[y]);
	return res;
}
signed main(){
	cin>>n;
	for(int i=1,u,v;i<n;i++){
		cin>>u>>v;
		add(u,v);
		add(v,u);
	}
	for(int i=1;i<=n;i++)
	cin>>w[i];
	dfs1(1,0),dfs2(1,1);
	build(1,1,order);
	cin>>q;
	for(int i=1,u,v;i<=q;i++){
		cin>>s;
		cin>>u>>v;
		if(s[0]=='C')change(1,seg[u],v);
		if(s[1]=='M')cout<<askmax1(u,v)<<endl;
		if(s[1]=='S')cout<<asksum1(u,v)<<endl;
	}
	return 0;
}
2023/1/8 10:18
加载中...