MnZn刚学树剖1ms,板子10pts求助
查看原帖
MnZn刚学树剖1ms,板子10pts求助
526895
WYZ20030051楼主2023/8/17 16:32

rt,AC on #4

#include<iostream>
#include<cstdio>
#include<cmath>
#include<string>
#include<cstring>
#include<algorithm>
#include<cassert>
#include<stack>
#include<queue>
#include<vector>
#include<map>
#include<cstdlib>
using namespace std;
#define ll long long
#define ull unsigned long long
int read()
{
	int now=0,nev=1; 
	char c=getchar();
	while(c<'0' || c>'9') 
	{ 
		if(c=='-') 
			nev=-1; 
		c=getchar();
	}
	while(c>='0' && c<='9') 
	{ 
		now=(now<<1)+(now<<3)+(c&15); 
		c=getchar(); 
	}
	return now*nev;
}
const int MAXN=1e5+10;
const int INF=1e18;
int n,q;
int w[MAXN];
int head[MAXN],tail[MAXN],nxt[MAXN],tt=0;
void add(int x,int y)
{
	nxt[++tt]=head[x];
	head[x]=tt;
	tail[tt]=y;
}
int dep[MAXN],fa[MAXN],son[MAXN];
int dfn[MAXN],seg[MAXN],size[MAXN],top[MAXN];
int cnt=0;
void dfs1(int u,int f)
{
	size[u]=1;
	for(int i=head[u];i;i=nxt[i])
	{
		int v=tail[i];
		if(v==f)
			continue;
		dep[v]=dep[u]+1;
		fa[v]=u;
		dfs1(v,u);
		size[u]+=size[v];
		if(size[v]>size[son[u]])
			son[u]=v;
	}
}
void dfs2(int u,int tp)
{
	dfn[u]=++cnt;
	seg[cnt]=u;
	top[u]=tp;
	if(son[u])
		dfs2(son[u],tp);
	for(int i=head[u];i;i=nxt[i])
	{
		int v=tail[i];
		if(v==fa[u] || v==son[u])
			continue;
		dfs2(v,v);
	}
}
struct node
{
	int l,r;
	ll sum,maxx;
}tr[MAXN<<2];
void build(int k,int l,int r)
{
	tr[k].l=l;
	tr[k].r=r;
	if(tr[k].l==tr[k].r)
	{
		tr[k].maxx=w[seg[l]];
		tr[k].sum=w[seg[l]];
		return ;
	}
	int m=l+r>>1;
	build(k<<1,l,m);
	build(k<<1|1,m+1,r);
	tr[k].maxx=max(tr[k<<1].maxx,tr[k<<1|1].maxx);
	tr[k].sum=tr[k<<1].sum+tr[k<<1|1].sum;
}
void modefy(int k,int l,int r,int x,int v)
{
	if(x<l || x>r)
		return ;
	if(x==l && x==r)
	{
		tr[k].sum=tr[k].maxx=v;
		return ;
	}
	int m=l+r>>1;
	if(x<=m)
		modefy(k<<1,l,m,x,v);
	if(x>m)
		modefy(k<<1|1,m+1,r,x,v);
	tr[k].maxx=max(tr[k<<1].maxx,tr[k<<1|1].maxx);
	tr[k].sum=tr[k<<1].sum+tr[k<<1|1].sum;
}
ll querymax(int k,int l,int r,int x,int y)
{
	if(x<=l && r<=y)
		return tr[k].maxx;
	int m=l+r>>1;
	ll res=-0x3f;
	if(x<=m)
		res=max(res,querymax(k<<1,l,m,x,y));
	if(y>m)
		res=max(res,querymax(k<<1|1,m+1,r,x,y));
	return res;
}
ll querysum(int k,int l,int r,int x,int y)
{
	if(x<=l && r<=y)
		return tr[k].sum;
	int m=l+r>>1;
	ll res=0;
	if(x<=m)
		res+=querysum(k<<1,l,m,x,y);
	if(y>m)
		res+=querysum(k<<1|1,m+1,r,x,y);
	return res;
}
ll askmax(int u,int v)
{
	ll res=-INF;
	while(top[u]!=top[v])
	{
		if(dep[top[u]]<dep[top[v]])
			swap(u,v);
		res=max(res,querymax(1,1,n,dfn[top[u]],dfn[u]));
		u=fa[top[u]];
	}
	if(dep[u]<dep[v])
		swap(u,v);
	res=max(res,querymax(1,1,n,dfn[v],dfn[u]));
	return res;
}
ll asksum(int u,int v)
{
	ll res=0;
	while(top[u]!=top[v])
	{
		if(dep[top[u]]<dep[top[v]])
			swap(u,v);
		res+=querysum(1,1,n,dfn[top[u]],dfn[u]);
		u=fa[top[u]]; 
	}
	if(dep[u]<dep[v])
		swap(u,v);
	res+=querysum(1,1,n,dfn[v],dfn[u]);
	return res;
}
int main()
{
	n=read();
	memset(head,0,sizeof(head));
	for(int i=1;i<=n-1;i++)
	{
		int x,y;
		x=read(),y=read();
		add(x,y);
		add(y,x);
	}
	for(int i=1;i<=n;i++)
		w[i]=read();
	dfs1(1,0);
	dfs2(1,0);
	build(1,1,n);
	q=read();
	for(int i=1;i<=q;i++)
	{
		string op;
		cin>>op;
		if(op[1]=='H')
		{
			int u,t;
			u=read(),t=read();
			modefy(1,1,n,dfn[u],t);
		}
		if(op[1]=='M')
		{
			int u,v;
			u=read(),v=read();
			printf("%lld\n",askmax(u,v));
		}
		if(op[1]=='S')
		{
			int u,v;
			u=read(),v=read();
			printf("%lld\n",asksum(u,v));
		}
	}
	return 0;
}
2023/8/17 16:32
加载中...