求助,大佬! 只能过#4
查看原帖
求助,大佬! 只能过#4
649246
nothing__楼主2023/6/15 20:55

全wa

#include<bits/stdc++.h>
using namespace std;
const int N=3e5+10;
int n, m, tsp, q[N], rnk[N];
struct edge{int x, y, pre;} a[N<<1]; int alen, last[N];
void ins(int x, int y) {a[++alen]={x, y, last[x]}; last[x]=alen;}
struct node{int fa, son, size, top, dep, dfn;} t[N]; 
struct trnode{int l, r, lc, rc, c, mx, tag;} tr[N]; int trlen;
void dfs1(int x, int fa, int d) {
	t[x].dep=d; t[x].fa=fa; t[x].size=1;
	for(int k=last[x];k;k=a[k].pre) {
		int y=a[k].y; if(y==fa) continue;
		dfs1(y, x, d+1); t[x].size+=t[y].size;
		if(t[y].size>t[t[x].son].size) t[x].son=y;
	}
}
void dfs2(int x, int top) {
	t[x].top=top; t[x].dfn=++tsp; rnk[tsp]=x;
	if(!t[x].son) return ; dfs2(t[x].son, top);
	for(int k=last[x],y;k;k=a[k].pre) 
		if((y=a[k].y)!=t[x].son&&y!=t[x].fa) dfs2(y, y);
}
#define ls(x) tr[x].lc 
#define rs(x) tr[x].rc 
#define size(x) (tr[x].r-tr[x].l+1)
void push_up(int x) {
	tr[x].c=(tr[ls(x)].c+tr[rs(x)].c);
	tr[x].mx=max(tr[ls(x)].mx, tr[rs(x)].mx);
}
void build(int l, int r) {
	int now=++trlen;
	tr[now]={l, r, -1, -1, 0, 0, 0};
	if(l==r) {tr[now].c=tr[now].mx=q[rnk[l]]; return ;}
	int mid=(l+r) >> 1;
	tr[now].lc=trlen+1, build(l, mid);
	tr[now].rc=trlen+1, build(mid+1, r);
	tr[now].c=(tr[tr[now].lc].c+tr[tr[now].rc].c);
	tr[now].mx=max(tr[tr[now].lc].mx, tr[tr[now].rc].mx);
}
void change(int now, int l, int r, int x, int c) {
	if(l==r) {tr[now].c=tr[now].mx=c; return ;}
	int mid=(tr[now].l+tr[now].r)>>1;
	if(x<=mid) change(ls(now), l, mid, x, c);
	else change(rs(now), mid+1, r, x, c);
	push_up(now);
}
int solve_mx(int now, int l, int r, int p, int q) {
	if(l>=p&&r<=q) return tr[now].mx;
	int mid=(l+r)>>1, mx=0;
	if(p<=mid) mx=max(mx,solve_mx(ls(now), l, mid, p, q));
	if(q>mid) mx=max(mx, solve_mx(rs(now), mid+1, r, p, q));
	return mx;
}
int solve_sum(int now, int l, int r, int p, int q) {
	if(l>=p&&r<=q) return tr[now].c;
	int mid=(l+r)>>1, sum=0;
	if(p<=mid) sum+=solve_sum(ls(now), l, mid, p, q);
	if(q>mid) sum+=solve_sum(rs(now), mid+1, r, p, q);
	return sum;
}
int ask_mx(int x, int y) {
	int mx=0;
	while(t[x].top!=t[y].top) {
		if(t[t[x].top].dep<t[t[y].top].dep) swap(x, y);
		mx=max(mx, solve_mx(1, 1, n, t[t[x].top].dfn, t[x].dfn));
		x=t[t[x].top].fa;
	}
	if(t[x].dfn>t[y].dfn) swap(x, y);
	mx=max(mx, solve_mx(1, 1, n, t[x].dfn, t[y].dfn));
	return mx;
}
int ask_sum(int x, int y) {
	int res=0;
	while(t[x].top!=t[y].top) {
		if(t[t[x].top].dep<t[t[y].top].dep) swap(x, y);
		res+=solve_sum(1, 1, n, t[t[x].top].dfn, t[x].dfn); x=t[t[x].top].fa;
	}
	if(t[x].dfn>t[y].dfn) swap(x, y);
	res+=solve_sum(1, 1, n, t[x].dfn, t[y].dfn);
	return res;
}
int main() {
	scanf("%d", &n);
	alen=0; memset(last, 0, sizeof(last));
	for(int i=1;i<n;i++) {
		int x, y; scanf("%d%d", &x, &y);
		ins(x, y); ins(y, x);
	}
	for(int i=1;i<=n;i++) scanf("%d", &q[i]);
	dfs1(1, 0, 1); dfs2(1, 1); build(1, n);
	scanf("%d", &m);
	while(m--) {
		char ss[10]; cin >> ss;
		int x, y; scanf("%d%d", &x, &y);
		if(ss[1]=='H') change(1, 1, n, t[x].dfn, y);
		else if(ss[1]=='M') printf("%d\n", ask_mx(x, y));
		else printf("%d\n", ask_sum(x, y));
	}
	return 0;
}
2023/6/15 20:55
加载中...