做题笔记 – P4242 树上的毒瘤

ooliver 发布于 4 小时前 41 次阅读 OI


AI 摘要

都说这题是毒瘤,5KB代码写得人麻了。可当你看穿“LCA处颜色相等”的那一眼,一切瞬间通透——颜色段和减1,虚树边权直接搞定。点分治、树剖全是为这个巧思服务的。来,一起拆解这棵“毒瘤”。

P4242 树上的毒瘤

这题太累人了,就算是我这种马蜂紧凑的也爆写了将近 5 KB 的代码。

观察一下题目,先是维护树上颜色段,想到树剖+线段树;再是若干关键点之间两两求贡献,想到虚树+点分治。写点分治时,我们固定一个 LCA,不难发现两条链合并时 LCA 处颜色相等,那贡献就是颜色段和减 $1$。那我们建虚树时就直接把两个关键点在原图上的链的颜色段数减一作为边权,然后统计答案时就加上一个 $1$ 即可。

代码:

C++
#include<bits/stdc++.h>
using namespace std;
#define ls p<<1
#define rs p<<1|1

const int N=1e5+5;
int n,q,op,a,b,c,idx,sum,rt,tot,sumd;
int col[N],dot[N],nd[N],tag[N];
int fa[N],dep[N],dfn[N],top[N],son[N],inv[N];
int sz[N],vis[N],ans[N],dp[N],siz[N],pdis[N],sumv[N];
vector<int> g[N],sonv[N];
vector<pair<int,int>> vt[N];
struct node{
    int sum,lc,rc,tag;
}tr[N<<2];

bool cmp(int x,int y){
    return dfn[x]<dfn[y];
}

void dfs1(int u,int f){
    dep[u]=dep[f]+1,fa[u]=f,siz[u]=1;
    for(int v:g[u]){
        if(v==f) continue;
        dfs1(v,u);
        siz[u]+=siz[v];
        if(siz[v]>siz[son[u]]) son[u]=v;
    }
}

void dfs2(int u,int tf){
    dfn[u]=++idx,inv[idx]=u,top[u]=tf;
    if(son[u]) dfs2(son[u],tf);
    for(int v:g[u]){
        if(v==fa[u]||v==son[u]) continue;
        dfs2(v,v);
    }
}

int lca(int x,int y){
    while(top[x]!=top[y]){
        if(dep[top[x]]<dep[top[y]]) swap(x,y);
        x=fa[top[x]];
    }
    return dep[x]<dep[y]?x:y;
}

node merge(node x,node y){
    if(!x.sum) return y;
    return {x.sum+y.sum-(x.rc==y.lc),x.lc,y.rc,0};
}

void pushup(int p){
    tr[p]=merge(tr[ls],tr[rs]);
}

void pushdown(int p){
    if(tr[p].tag){
        tr[ls]={1,tr[p].tag,tr[p].tag,tr[p].tag};
        tr[rs]={1,tr[p].tag,tr[p].tag,tr[p].tag};
        tr[p].tag=0;
    }
}

void build(int p,int l,int r){
    if(l==r){
        tr[p]={1,col[inv[l]],col[inv[l]],0};
        return;
    }
    int mid=(l+r)>>1;
    build(ls,l,mid);
    build(rs,mid+1,r);
    pushup(p);
}

void change(int p,int l,int r,int x,int y,int k){
    if(l>=x&&r<=y){
        tr[p]={1,k,k,k};
        return;
    }
    int mid=(l+r)>>1;
    pushdown(p);
    if(x<=mid) change(ls,l,mid,x,y,k);
    if(y>mid) change(rs,mid+1,r,x,y,k);
    pushup(p);
}

node get(int p,int l,int r,int x,int y){
    if(l>=x&&r<=y) return tr[p];
    int mid=(l+r)>>1;
    pushdown(p);
    if(y<=mid) return get(ls,l,mid,x,y);
    if(x>mid) return get(rs,mid+1,r,x,y);
    return merge(get(ls,l,mid,x,y),get(rs,mid+1,r,x,y));
}

void pchange(int x,int y,int k){
    while(top[x]!=top[y]){
        if(dep[top[x]]<dep[top[y]]) swap(x,y);
        change(1,1,n,dfn[top[x]],dfn[x],k);
        x=fa[top[x]];
    }
    if(dep[x]>dep[y]) swap(x,y);
    change(1,1,n,dfn[x],dfn[y],k);
}

int pquery(int x,int y){
    node res={0,0,0,0};
    while(top[x]!=top[y]){
        res=merge(get(1,1,n,dfn[top[x]],dfn[x]),res);
        x=fa[top[x]];
    }
    res=merge(get(1,1,n,dfn[y],dfn[x]),res);
    return res.sum;
}

void buildvt(){
    sort(dot+1,dot+1+a,cmp);
    for(int i=1;i<=a;i++){
        nd[++sum]=dot[i];
        if(i<a) nd[++sum]=lca(dot[i],dot[i+1]);
    }
    sort(nd+1,nd+1+sum,cmp);
    sum=unique(nd+1,nd+1+sum)-nd-1;
    for(int i=1;i<sum;i++){
        int f=lca(nd[i],nd[i+1]),wis=pquery(nd[i+1],f);
        vt[f].push_back({nd[i+1],wis-1});
        vt[nd[i+1]].push_back({f,wis-1});
    }
}

void clearvt(){
    for(int i=1;i<sum;i++){
        int f=lca(nd[i],nd[i+1]);
        vt[f].clear();
        vt[nd[i+1]].clear();
        vis[f]=vis[nd[i+1]]=ans[f]=ans[nd[i+1]]=0;
    }
    sum=tot=0;
}

void getrt(int u,int f){
    siz[u]=1,dp[u]=0;
    for(auto[v,w]:vt[u]){
        if(v==f||vis[v]) continue;
        getrt(v,u);
        siz[u]+=siz[v];
        dp[u]=max(dp[u],siz[v]);
    }
    dp[u]=max(dp[u],tot-siz[u]);
    if(dp[u]<dp[rt]) rt=u;
}

void getdis(int u,int f,int dis,int rot){
    if(tag[u]){
        pdis[u]=dis,sumd+=dis;
        sonv[rot].push_back(u);
    }
    for(auto[v,w]:vt[u]){
        if(v==f||vis[v]) continue;
        getdis(v,u,dis+w,rot);
    }
}

void calc(int u){
    int totd=0;
    sumd=0;
    for(auto[v,w]:vt[u]){
        if(vis[v]) continue;
        sumv[v]=sumd;
        getdis(v,u,w,v);
        sumv[v]=sumd-sumv[v],totd+=sonv[v].size();
    }
    for(auto[v,w]:vt[u]){
        if(vis[v]) continue;
        for(auto x:sonv[v]){
            ans[x]+=(pdis[x]+1)*(totd-sonv[v].size())+(sumd-sumv[v]);
            if(tag[u]) ans[x]+=(pdis[x]+1),ans[u]+=(pdis[x]+1);
        }
        sonv[v].clear();
    }
}

void solve(int u){
    vis[u]=1;
    calc(u);
    for(auto[v,w]:vt[u]){
        if(vis[v]) continue;
        tot=siz[v],dp[v]=1e9,rt=0;
        getrt(v,0);
        solve(rt);
    }
}

int main(){
    cin>>n>>q;
    for(int i=1;i<=n;i++) cin>>col[i];
    for(int i=1,u,v;i<n;i++){
        cin>>u>>v;
        g[u].push_back(v);
        g[v].push_back(u);
    }
    dfs1(1,0);
    dfs2(1,1);
    build(1,1,n);
    while(q--){
        cin>>op;
        if(op==1){
            cin>>a>>b>>c;
            pchange(a,b,c);
        }
        else{
            cin>>a;
            for(int i=1;i<=a;i++) cin>>dot[i],tag[dot[i]]=1;
            buildvt();
            getrt(nd[1],0);
            solve(nd[1]);
            for(int i=1;i<=a;i++){
                cout<<ans[dot[i]]<<" ";
                tag[dot[i]]=0;
            }
            cout<<"\n";
            clearvt();
        }
    }
    return 0;
}