求助!样例过但是 0 WA
查看原帖
求助!样例过但是 0 WA
744354
Wil_Lei楼主2023/7/14 11:27
#include <bits/stdc++.h>
using namespace std;
const int N=3e4+10;
struct Segment_Tree {
	int l,r,s[2],sum,mx;
} tr[N<<2];
int n,q,fa[N],w[N],dfn[N],pos[N],dfc;
int sz[N],big[N],dep[N],top[N];
vector<int> g[N];char op[10];
void dfs1(int u,int pa) {
	sz[u]=1,dep[u]=dep[pa]+1;
	for (int v:g[u]) {
		if (v==pa) continue;
		dfs1(v,u),fa[v]=u,sz[u]+=sz[v];
		if (sz[big[u]]<sz[v]) big[u]=v;
	}
}
void dfs2(int u,int tp) {
	if (!u) return;
	dfn[++dfc]=u,pos[u]=dfc;
	top[u]=tp,dfs2(big[u],tp);
	for (int v:g[u]) {
		if (v==fa[u] || v==big[u]) continue;
		dfs2(v,v);
	}
}//树链剖分
void pushup(int u) {
	int ls=tr[u].s[0],rs=tr[u].s[1];
	tr[u].mx=max(tr[ls].mx,tr[rs].mx);
	tr[u].sum=tr[ls].sum+tr[rs].sum;
}
void build(int u,int l,int r) {
	if (l==r) {
		int wt=w[pos[l]];
		tr[u]={l,r,{0,0},wt,wt};
		return;
	}
	int mid=(l+r)>>1;
	int ls=(u<<1),rs=(u<<1|1);
	tr[u]={l,r,{ls,rs},0,0};
	build(ls,l,mid),build(rs,mid+1,r);
	pushup(u);
}
void update(int u,int k,int t) {
	if (!tr[u].s[0]) {
		tr[u].mx=tr[u].sum=t;
		return;
	}
	int mid=tr[tr[u].s[0]].r;
	int ls=tr[u].s[0],rs=tr[u].s[1];
	if (pos[k]<=mid) update(ls,k,t);
	else update(rs,k,t);
	pushup(u);
}
int qmax(int u,int l,int r) {
	if (l<=tr[u].l && tr[u].r<=r)
		return tr[u].mx;
	int mid=tr[tr[u].s[0]].r,a=0,b=0;
	if (l<=mid) a=qmax(tr[u].s[0],l,r);
	if (mid<r) b=qmax(tr[u].s[1],l,r);
	return max(a,b);
}
int qsum(int u,int l,int r) {
	if (l<=tr[u].l && tr[u].r<=r)
		return tr[u].sum;
	int mid=tr[tr[u].s[0]].r,a=0,b=0;
	if (l<=mid) a=qsum(tr[u].s[0],l,r);
	if (mid<r) b=qsum(tr[u].s[1],l,r);
	return a+b;
}//线段树
int main() {
	scanf("%d",&n);
	for (int i=1,a,b; i<n; i++) {
		scanf("%d%d",&a,&b);
		g[a].push_back(b);
		g[b].push_back(a);
	}
	for (int i=1; i<=n; i++)
		scanf("%d",w+i);
	dfs1(1,0),dfs2(1,1);
	build(1,1,n);
	scanf("%d",&q);
	for (int i=0,u,v; i<q; i++) {
		scanf("%s%d%d",op+1,&u,&v);
		if (op[1]=='C') update(1,u,v);
		else if (op[4]=='X') {
			int ans=-0x3f3f3f3f;
			while (top[u]!=top[v]){
				if (dep[top[u]]<dep[top[v]])
					swap(u,v);
				int qry=qmax(1,pos[top[u]],pos[u]);
				ans=max(ans,qry),u=fa[top[u]];
			}
			if (dep[u]>dep[v]) swap(u,v);
			ans=max(ans,qmax(1,pos[u],pos[v]));
			printf("%d\n",ans);
		}else{
			int ans=0;
			while (top[u]!=top[v]){
				if (dep[top[u]]<dep[top[v]])
					swap(u,v);
				int qry=qsum(1,pos[top[u]],pos[u]);
				ans+=qry,u=fa[top[u]];
			}
			if (dep[u]>dep[v]) swap(u,v);
			ans+=qsum(1,pos[u],pos[v]);
			printf("%d\n",ans);
		}
	}
	return 0;
}
2023/7/14 11:27
加载中...