萌新刚学树剖,只过了样例
查看原帖
萌新刚学树剖,只过了样例
600112
_HHJ楼主2022/9/16 17:01
#include <iostream>
#define lson ls,l,mid 
#define rson rs,mid+1,r
#define ls p<<1
#define rs p<<1|1
#define mid (l+r>>1)
#define int long long 
using namespace std;
const int N = 2e6+10;
int h[N],ne[N<<1],e[N<<1],idx;
int w[N],nw[N],son[N],sz[N],cnt;
int id[N],dep[N],fa[N],top[N];
int sum[N<<2], maxn[N<<2];
int n,q;    char opt[20];
inline int read()
{
    int res = 0, flag = 1;
    char ch = getchar();
    while(ch < '0'|| ch > '9')
    {
        if(ch == '-')
        flag = -1 ;
        ch = getchar();
    }
    while(ch >= '0'&&ch <= '9')
    {
        res = (res << 1) + (res << 3) + (ch^48);
        ch = getchar();
    }
    return res*flag;
}
void add(int u,int v){
    e[++idx] = v;
    ne[idx] = h[u];
    h[u] = idx;
}
void dfs1(int u,int father){
    sz[u] = 1,dep[u] = dep[father] + 1,fa[u] = father;
    for (int i = h[u] ; i ; i = ne[i]){
        int v = e[i];
        if(v == father)continue;
        dfs1(v,u);
        sz[u] += sz[v];
        if(sz[son[u]] < sz[v])
        son[u] = v;
    }
}
void dfs2(int u,int t){
    top[u] = t,id[u] = ++cnt,nw[cnt] = w[u];
    if(!son[u])return;
    dfs2(son[u],t);
    for (int i = h[u] ; i ; i = ne[i]){
        int v = e[i];
        if(v == fa[u]|| v == son[u])continue;
            dfs2(v,v);
    }


}
void update(int p){
    sum[p] = sum[ls] + sum[rs];
    maxn[p] = max(maxn[ls],maxn[rs]);

}
void build(int p,int l,int r){
    if(l == r){
        sum[p] = maxn[p] = nw[l];
        return;
    }
    build(lson); build(rson);
    update(p);
}
void change(int p,int l,int r,int u,int t){
    if(l == r) {
        sum[p] =  maxn[p] = t; 
        return ;
    }
    if(u <= mid)change(lson,u,t);
    else change(rson,u,t);
    update(p);
}
int ask_sum(int p,int l,int r,int x,int y){
    if(x <= l && r <= y) return sum[p];
    int res = 0;
    if(x <= mid)res += ask_sum(lson,x,y);
    if(y >  mid)res += ask_sum(rson,x,y);
    update(p);
    return res;
}
int ask_max(int p,int l,int r,int x,int y){
    if(x <= l && r <= y) return maxn[p];
    int res = -(1e9);
    if(x <= mid)res = max(ask_max(lson,x,y),res);
    if(y >  mid)res = max(ask_max(rson,x,y),res);
    update(p);
    return res;
}
int query_max(int u,int v){
    int res = -(1e9);
    while (top[u]!=top[v])
    {
        if(dep[top[u]] < dep[top[v]])swap(u,v);
        res = max(ask_max(1,1,n,id[top[u]],id[u]),res);
        u  = fa[top[u]];
    }
    if(dep[u] < dep[v])swap(u,v);
    res = max(res,ask_max(1,1,n,id[v],id[u]));
    return res;
    
}
int query_sum(int u,int v){
    int res = 0;
    while (top[u]!=top[v])  
    {
        if(dep[top[u]] < dep[top[v]])swap(u,v);
        res += ask_sum(1,1,n,id[top[u]],id[u]);
        u = fa[top[u]];
    }
    if(dep[u] < dep[v])swap(u,v);
    res += ask_sum(1,1,n,id[v],id[u]);
    return res;
    
}
signed main()
{
     //freopen("1.in","r",stdin);
     //freopen("1.out","w",stdout);
     n = read();
     for (int u,v,i = 1 ; i < n ; i ++ )
     {
        u = read();
        v = read();
        add(u,v);
        add(v,u);
     }
     for (int i = 1 ; i <= n ; i ++ )w[i] = read();
     dfs1(1,1);
     dfs2(1,1);
     build(1,1,n);
     q = read();
     while(q--){
        int u,v;
        scanf("%s",opt);
        u = read();
        v = read();
        if(opt[0] == 'Q'){
            if(opt[1] == 'M')
            printf("%lld\n",query_max(u,v));
            if(opt[1] == 'S')
            printf("%lld\n",query_sum(u,v));
        }else {
            change(1,1,n,id[u],v);
        }
     }
}
2022/9/16 17:01
加载中...