树上启发式合并

简介

树上启发式合并,主要用于解决要求每个子树内信息的一种暴力优化离线算法,将 \(O(n^2)\) 的暴力(树形 DP 形成 \(O(n)\),再清空又是 \(O(n)\),总计 \(O(n^2)\)),转换为了 \(O(n \log n)\) 的算法。主要思想是暴力枚举轻儿子所在子树的答案时清空残留数据,重儿子最后处理,使得重儿子的答案无需清空。

例题一

CF600E Lomsat gelral
给定一棵树,每个节点有一个颜色。求以每个节点为根的子树内节点颜色出现次数最多的颜色编号之和。

解法

树上启发式合并模板题。
我们可以用 dfs,首先,第一次 dfs 求出每个节点的重儿子和子树大小。代码如下:

void dfs1(int x,int fa)
{
    sz[x]=1;
    int mx=0;
    for(auto &i:G[x])
    {
        if(i!=fa)
        {
            dfs1(i,x);
            sz[x]+=sz[i];
            if(sz[i]>mx) son[x]=i,mx=sz[i];
        }
    }
}

然后,我们暴力求解一个子树内的答案。实现可以记录当前的答案 \(ans\) 和最大的出现次数 \(mx\)。注意,此时就需要跳过一开始传入的节点的重儿子代码如下:

void dfs2(int x,int fa,int p)
{
    cut[a[x]]++;
    if(cut[a[x]]>ma)
    {
        ma=cut[a[x]];
        maa=a[x];
    }
    else if(cut[a[x]]==ma) maa+=a[x];
    for(auto &i:G[x])
    {
        if(i==fa||i==p) continue;
        dfs2(i,x,p);//注意,p不变
    }
}

接下来,由于我们需要每一个点的子树答案,所以我们需要清空数据。代码如下:

void init(int x,int fa)
{
    cut[a[x]]--;
    for(auto &i:G[x])
        if(i!=fa) init(i,x);
}

下面是树上启发式合并最核心的部分。还是选择暴力的方式,轻儿子显然用上面的dfs2函数解决即可。而重儿子则可以直接将信息传给父亲,需要最后处理。这也就是我们在上面的dfs2中可以跳过一开始的重儿子的原因。代码如下:

void dfs3(int x,int fa)
{
    for(auto &i:G[x])
    {
        if(i!=son[x]&&i!=fa)
        {
            dfs3(i,x);
            init(i,x);
            ma=maa=0;
        }
    }
    if(son[x]) dfs3(son[x],x);
    dfs2(x,fa,son[x]);
    ans[x]=maa;
}

完整代码如下:

#include<bits/stdc++.h>
using namespace std;
#define int long long
const int MAXN=1e5+10;
int a[MAXN],n,sz[MAXN],son[MAXN],ans[MAXN],cut[MAXN],ma,maa;
vector<int> G[MAXN];
void dfs1(int x,int fa)
{
    sz[x]=1;
    int mx=0;
    for(auto &i:G[x])
    {
        if(i!=fa)
        {
            dfs1(i,x);
            sz[x]+=sz[i];
            if(sz[i]>mx) son[x]=i,mx=sz[i];
        }
    }
}
void dfs2(int x,int fa,int p)
{
    cut[a[x]]++;
    if(cut[a[x]]>ma)
    {
        ma=cut[a[x]];
        maa=a[x];
    }
    else if(cut[a[x]]==ma) maa+=a[x];
    for(auto &i:G[x])
    {
        if(i==fa||i==p) continue;
        dfs2(i,x,p);
    }
}
void init(int x,int fa)
{
    cut[a[x]]--;
    for(auto &i:G[x])
        if(i!=fa) init(i,x);
}
void dfs3(int x,int fa)
{
    for(auto &i:G[x])
    {
        if(i!=son[x]&&i!=fa)
        {
            dfs3(i,x);
            init(i,x);
            ma=maa=0;
        }
    }
    if(son[x]) dfs3(son[x],x);
    dfs2(x,fa,son[x]);
    ans[x]=maa;
}
signed main(  )
{
    cin>>n;
    for(int i=1;i<=n;i++) scanf("%lld",&a[i]);
    for(int i=1;i<n;i++)
    {
        int u,v;
        scanf("%lld%lld",&u,&v);
        G[u].push_back(v);
        G[v].push_back(u);
    }
    dfs1(1,0);dfs3(1,0);
    for(int i=1;i<=n;i++) cout<<ans[i]<<' ';
}

例题二

CF246E Blood Cousins Return
给定一森林,每个节点有字符串权值。有 m 次询问,每次询问给定 v 和 k,求节点 v 的所有 k 级儿子中,有多少不同的权值。

解法

注意到询问子树内信息且离线,自然想到树上启发式合并。我们可以预处理每个节点在其所在的树内的深度,然后将询问离线。设 \(dep_i\) 表示节点 \(u\) 的深度,可以发现问题是询问 \(v\) 子树内深度为 \(dep_v+k\) 的节点的答案。可以用 vector 存储询问,而答案可以很方便地用 set 维护。
set 维护每个深度的答案,然后用以上思路写出代码即可。完整代码如下:

#include<bits/stdc++.h>
using namespace std;
const int MAXN=1e5+10;
string a[MAXN];
int n,m,dep[MAXN],son[MAXN],sz[MAXN],ans[MAXN];
bool vis[MAXN];
vector<int> G[MAXN];
vector<pair<int,int> >q[MAXN];
void dfs1(int x)
{
    sz[x]=1;
    for(auto &i:G[x])
    {
        dep[i]=dep[x]+1;
        dfs1(i);
        sz[x]+=sz[i];
        if(sz[i]>sz[son[x]]) son[x]=i;
    }
}
set<string> S[MAXN*2];
void del(int x)//清空整个子树
{
    S[dep[x]].clear();
    for(auto &i:G[x]) del(i);
}
void add(int x)//加入整个子树
{
    S[dep[x]].insert(a[x]);
    for(auto &i:G[x]) add(i);
}
void dfs2(int x)//启发式合并部分
{
    vis[x]=true;
    for(auto &i:G[x]) if(i!=son[x]) dfs2(i),del(i);//求解、清空
    if(son[x]) dfs2(son[x]);
    for(auto &i:G[x]) if(i!=son[x]) add(i);
    S[dep[x]].insert(a[x]);
    for(auto &i:q[x]) ans[i.second]=S[i.first].size();
}
int main(  )
{
    ios::sync_with_stdio(false);
    cin.tie(0);
    cin>>n;
    for(int i=1;i<=n;i++)
    {
        int x;
        cin>>a[i]>>x;
        if(x) G[x].push_back(i);
    }
    for(int i=1;i<=n;i++) if(!dep[i]) dep[i]=1,dfs1(i);
    cin>>m;
    for(int i=1;i<=m;i++)
    {
        int x,y;
        cin>>x>>y;
        q[x].push_back({dep[x]+y,i});
    }
    for(int i=1;i<=n;i++) if(!vis[i]) dfs2(i),del(i);
    for(int i=1;i<=m;i++) cout<<ans[i]<<'\n';
}

持续更新ing...

posted @ 2026-05-16 15:59  duchuyuanX  阅读(25)  评论(0)    收藏  举报