学习心得 - 数据结构 - 虚树

· · 算法·理论

定义

虚树主要用于优化树形 dp。

每一次我们树形 dp 需要使用 k(\sum k \leq 10^5) 个点,设他们为关键节点。

我们的树是 n(\leq 10^5) 个点的。

每一次查询 (\leq 10^5),要是我们直接 dp,那么复杂度为 \mathcal O(nq),会超时。

考虑像可持久化一样,只在涉及的点上修改,复杂度为 \mathcal O(k)。

虚树建立

考虑涉及到的点是哪些。

对于一般的树形 dp,我们一般需要将操作移到 \operatorname{LCA} 上。

所以,我们可以考虑求出关键点两两之间的 \operatorname{LCA}。

那么,我们的复杂度退化成了 \mathcal O(k^2\log_2n)。

怎么办?

LCA 做法

我们发现,这些关键点的 \operatorname{LCA} 很多都是重复的。

一个朴素的想法是选择两个点,将他们合并到 \operatorname{LCA} 上,然后就可以 \mathcal O(k\log_2n)。

怎么合并?

选择按照 dfn 的顺序合并,就可以求出 k-1 个 \operatorname{LCA}。

我们一般是从根节点搜,所以我们可以加一个 1 点,虚树总点数是关键加 \operatorname{LCA} 再加根一共是 k+k-1+1=2k 个。

单调栈做法

我们可以把虚树切成一些链。

按照 dfn 序,不停插入链。

要是新加入的点不是栈的栈顶,那么就取出直到组成新链。

这样就可以用 \mathcal O(k\log_2n) 的复杂度建树了。

例题 1 【SX/NOI-】P2495 [SDOI2011] 消耗战

这里直接给了你 k 个关键点要求通过切断任意边让这 k 个点与 1 不相连。

我们所需边权可以直接 dfs 求点 1 到点 i 的路径上最小边权,设为 {\min}w_i。

先按照上述方法建造虚树。

然后考虑简单的 dp 就可以通过了。

#include<bits/stdc++.h>
using namespace std;
#define int long long
const int N=2.5e5+10,INF=1e18;
int n,m,minw[N],T,k[N];
int sta[N],top;
map<int,int>issp;
struct tnode{
    int dep,sz,top,id,fa,son;
}t[N];
vector<pair<int,int> >e[N],g[N];
bool cmp(int u,int v){
    return t[u].id<t[v].id;
}
int dfncnt;
void dfs1(int p,int fa){
    t[p].id=++dfncnt;
    t[p].sz=1;
    for(auto v:e[p])if(v.first!=fa){
        t[v.first].fa=p;
        t[v.first].dep=t[p].dep+1;
        minw[v.first]=min(minw[p],v.second);
        dfs1(v.first,p);
        t[p].sz+=t[v.first].sz;
        if(t[t[p].son].sz<t[v.first].sz)
            t[p].son=v.first;
    }
    return;
}
void dfs2(int p,int tp){
    t[p].top=tp;
    if(!t[p].son)
        return;
    dfs2(t[p].son,tp);
    for(auto v:e[p])if(v.first!=t[p].fa&&v.first!=t[p].son)
        dfs2(v.first,v.first);
}
int LCA(int u,int v){
    while(t[u].top!=t[v].top){
        if(t[t[u].top].dep<t[t[v].top].dep)
            swap(u,v);
        u=t[t[u].top].fa;
    }
    if(t[u].dep>t[v].dep)swap(u,v);
    return u;
}   
int dp(int u){
    if(g[u].size()==0)
        return minw[u];
    int su=0;
    for(auto v:g[u])
        su+=dp(v.first);
    g[u].clear();
    return min(minw[u],su);
}
signed main(){
    cin>>n;
    for(int i=1;i<n;i++){
        int u,v,w;
        cin>>u>>v>>w;
        e[u].push_back({v,w});
        e[v].push_back({u,w});
    }
    minw[1]=INF;
    dfs1(1,0),dfs2(1,1);
    cin>>m;
    while(m--){
        cin>>T;
        issp.clear();
        for(int i=1;i<=T;i++)
            cin>>k[i],issp[k[i]]=1;
        sort(k+1,k+T+1,cmp);
        top=0;
        sta[++top]=1;
        sta[++top]=k[1];
        for(int i=2;i<=T;i++){
            int lca=LCA(k[i],sta[top]);
            if(lca!=sta[top]){
                while(t[sta[top-1]].id>=t[lca].id)
                    g[sta[top-1]].push_back({sta[top],0}),
                    --top;
                if(sta[top]!=lca)
                    g[lca].push_back({sta[top],0}),
                    sta[top]=lca;
                sta[++top]=k[i];
            }else;
        }
        while(top)
            g[sta[top-1]].push_back({sta[top],0}),
            --top;
        cout<<dp(1)<<"\n";
    }
    return 0;
}