改了很久还是没有看出问题,下面是代码
#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;
}