不开 O2 MLE ,开 O2 TLE ,求助大佬
#include <iostream>
#include <vector>
#define mid ((l+r)>>1)
#define lson pos<<1
#define rson pos<<1|1
using namespace std;
const int inf = 1e9;
const int MAXN = 3e4+5;
int n, m, a[MAXN], dep[MAXN], fa[MAXN], siz[MAXN], mxson[MAXN];
struct Seg{
int mx, sum;
}tree[MAXN<<2];
int top[MAXN], idx, dfn[MAXN], id[MAXN];
vector <int> G[MAXN];
void dfs1(int u,int pre)
{
fa[u] = pre, dep[u] = dep[pre]+1;
siz[u] = 1;
for(auto v:G[u])
{
if(v==pre) continue;
dfs1(v,u);
siz[u] += siz[v];
if(siz[v]>siz[mxson[u]]) mxson[u] = v;
}
}
void dfs2(int u,int start)
{
dfn[u] = ++idx; id[idx] = u;
top[u] = start;
if(mxson[u]==0) return ;
dfs2(mxson[u], start);
for(auto v:G[u])
{
if(v==mxson[u] or v==fa[u]) continue;
dfs2(v,v);
}
}
void pushup(int pos)
{
tree[pos].sum = tree[lson].sum+tree[rson].sum;
tree[pos].mx = max(tree[lson].mx, tree[rson].mx);
}
void build(int pos,int l,int r)
{
if(l==r)
{
tree[pos].sum = tree[pos].mx = a[id[l]];
return ;
}
build(lson,l,mid);
build(rson,mid+1,r);
pushup(pos);
}
void modify(int pos,int l,int r,int x,int val)
{
if(l==r) {
tree[pos].sum = tree[pos].mx = val;
return ;
}
if(x<=mid) modify(lson,l,mid,x,val);
else modify(rson,mid+1,r,x,val);
pushup(pos);
}
Seg query(int pos,int l,int r,int x,int y)
{
if(x<=l && r<=y) return tree[pos];
else if(y<=mid) return query(lson,l,mid,x,y);
else if(x>mid) return query(rson,mid+1,r,x,y);
else {
Seg res1 = query(lson,l,mid,x,y), res2 = query(rson,mid+1,r,x,y), res;
res.mx = max(res1.mx, res2.mx);
res.sum = res1.sum + res2.sum;
return res;
}
}
Seg solve(int x,int y)
{
Seg ret; ret.mx = -inf, ret.sum=0;
int fx = top[x], fy = top[y];
while(fx!= fy)
{
if(dep[fx] > dep[fy])
{
Seg res = query(1,1,n,dfn[fx],x);
ret.sum += res.sum, ret.mx = max(ret.mx , res.mx);
x = fa[fx];
}
else
{
Seg res = query(1,1,n,dfn[fy],y);
ret.sum += res.sum, ret.mx = max(ret.mx , res.mx);
y = fa[fy];
}
fx = top[x], fy = top[y];
}
if(dfn[x] < dfn[y])
{
Seg res = query(1,1,n,dfn[x],dfn[y]);
ret.mx = max(ret.mx,res.mx);
ret.sum += res.sum;
}
else
{
Seg res = query(1,1,n,dfn[y],dfn[x]);
ret.mx = max(ret.mx,res.mx);
ret.sum += res.sum;
}
return ret;
}
int main()
{
cin >> n;
for(int i=1;i<n;i++)
{
int u,v; cin >> u >> v;
G[u].push_back(v); G[v].push_back(u);
}
for(int i=1;i<=n;i++)
cin >> a[i];
dfs1(1,0);
dfs2(1,1);
build(1,1,n);
cin >> m;
while(m--)
{
string opt; int x,y;
cin >> opt >> x >> y;
if(opt=="CHANGE") modify(1,1,n,x,y);
if(opt=="QMAX") cout << solve(x,y).mx << endl;
if(opt=="QSUM") cout << solve(x,y).sum << endl;
}
}