欧拉序+ST表求解LCA

· · 个人记录

欧拉序&ST表&LCA

这篇我自认为是比较有必要写的,之前的哈希,高斯消元等等都写到一半就鸽了,这次一定要写完,至于那些等回头再补吧(`・ω・´)

首先是这篇博客要讲的几个定义:

LCA

即最近公共祖先,对于一棵树T上的节点a,b,他们显然都有共同的若干个祖宗,最近公共祖先即他们的祖宗里辈分最小的。

比如我和强者wqx是堂兄弟(我竟然和dalao是亲戚!),那么我们的共同祖宗有很多,比如太祖父,曾祖父,但是辈分最小的是爷爷

那么给定你一棵树,树的大小为n,有m次询问,求

那么大家都知道,求LCA的方法主要为倍增和tarjan

倍增的复杂度为预处理O(n\log_2(n)),在线求解,每次询问复杂度O(\log_2(n)),即总复杂度为O(n\log_2(n)+m\log_2(n))已经可以满足基本要求,洛谷的那个题可以通过

tarjan的复杂度为O(\alpha (n)n),单次询问O(1)基本可以看做是线性的,但是强制离线。

今天的欧拉序列+ST表的复杂度为预处理O(n\log_2(n)+n),询问为O(1),所以基本上是可以完全淘汰倍增(在m非常小,n非常大的情况下或许倍增能占优,但是绝大多数情况肯定是欧拉序更好)

那么欧拉序列这个东西是什么?

欧拉序列

我们给出一个具体的树

首先大家都知道dfs序,这棵树的dfs序为

1 2 4 5 7 3 6 8

那么欧拉序,就是dfs回溯的时候也算上这个点,也就是说有些点可能会出现不止1次,欧拉序列如下:

1 2 4 2 5 7 5 2 1 3 6 8 6 3 1,我们暂且将这个数组命名为Euler

可以证明的是,这个长度为2n-1,证明在此就略过了

那么欧拉序列与LCA有何关系?

我们先定义depth数组,depth[i]表示Euler[i]的深度,则其为:

1 2 3 2 3 4 3 2 1 2 3 4 3 2 1

假设我们要求4和7的LCA,通过肉眼观察,我们可以知道是2,那么这与欧拉序列有啥关系?

我们将4第一次出现的位置找出来,将7第一次出现的位置找出来,把这之间的序列拿出来

Euler: 4 2 5 7

depth:3 2 3 4

我们可以发现,正好是这个序列中深度最低的点。为什么?

证明

下面是简要语言叙述证明(这部分是我自己脑补的,可能语言有些繁琐,大家只要知道这是对的就可以,有兴趣可以看一看)

证法 一

我们可以这样考虑,有两个节点a与b,有如下推论:

A:他们第一次出现在Euler序列的时候,一定是第一次被访问到,一定不是回溯。

B:通过A推论可知,如果是Euler[i]——Euler[i+k]的点都是第一次出现,那么其深度序列depth[i]——depth[i+k]必然为单调递增的序列。

C:由B推得,若其depth序列不为递增,那么一定是回溯到了某个点。若满足以下两个条件:

1,depth[i+n]是第一个深度小于depth[i]的

2,depth[i+n]——depth[i+n+k]为单调递减序列

则Euler[i+n]——Euler[i+n+k]都是i的祖宗。证明过程略,但这是成立的。我真的懒得想怎么证C了,大家知道对就好

假设Euler[i]=a,Euler[i-k]到Euler[i] (k未知,k>=0,且一定存在这样的k)必然为单调递增,都是a的祖宗。也有若干个点在i后并且满足推论C的条件,即在i后面深度比他小的点为其祖宗。假设Euler[j]=b,则j之前一定也有若干个点单调递增,则他们都是b的祖宗。则一定有一个点既在i后的递减序列,也在j前的递增序列里,那么他就是a,b的祖宗。那么如何证明他是最近的公共祖先?如果有一个点maxp,他不是a与b的最近公共祖宗,但是是a与b的公共祖宗。那么maxp并不会出现在Euler[i]——Euler[j]里。因为我们知道dfs是深度优先,有深度更深的点,就会先去遍历它,假设a与b的最近公共祖先是c,则搜完a后回溯到c,一定会去找b而不是回溯到maxp。maxp一定是b以后才回溯到,所以Euler[i]——Euler[j]不会出现。

所以,Euler[i]——Euler[j]中深度最小的点为其最近公共祖先,一定不是他们的其他祖先。

证法二

显然,从a访问到b时一定会经过最近公共祖先,且根据证法1中的结论,这个区间内不会有不是他们的最近公共祖先但是却是他们的公共祖先的点,因此深度最小的点就是他们的最近公共祖先。

因此我们可以得出这样一个结论:两个点的LCA就是在欧拉序列中两个点第一次出现的位置之间的序列中深度最小的点。

我们将这个定义拆分出来,可以发现有这样几个要点:

1,在欧拉序列中

2,两个点第一次出现的位置

3,这两个位置之间的序列

4,深度最小。

因此我们可以发现,我们只要定义Euler序列,然后记录他们第一次出现的位置和深度,然后找区间最小值就行了。

但是区间最小值怎么找?如果说暴力的话还不如倍增,所以这里我们用ST表来解决区间最小值。

ST表

ST表,用来解决RMQ问题(区间最值问题),时间复杂度为预处理O(n\log_2(n)),单次查询为O(1),是一种非常优秀的算法。他表现为一个二维数组st[i][j],表示以j为起点到第j+2^i这个区间的的最小值。那么如果要查找a与b之间的最小值,可以定义一个k,且k为不大于b-a的2的最大幂。然后将a——b的区间分解为a——2^k,2^k+1 —— b,两个区间,取这两个区间的最大值就行了。

那么如何预处理st表?二重循环即可,第一层循环i,第二层循环j就行了,但是要注意边界条件,这个在下面的代码里有具体显示。

下面是代码:

void dfs(int nowp,int dep){
    //cout<<nowp<<" ";
    now_ord++;
    first_dfs[nowp]=now_ord;
    dfs_ord[now_ord]=nowp;
    depth[now_ord]=dep;
    for(int i=first[nowp];i;i=nextt[i]){
        int v=to[i];
        if(!first_dfs[v]){
            dfs(v,dep+1);
            //cout<<nowp<<" ";
            now_ord++;
            dfs_ord[now_ord]=nowp,depth[now_ord]=dep;
        }
    }
}

这个是dfs部分,直接处理出Euler,first,depth三个数组,这里名字不太一样。(我尽量用顾名思义的名字)

dfs_ord[i]表示Euler序列的第i个(dfs的顺序,这名字没毛病)

depth[i]表示Euler序列的第i个元素的点的深度,注意,这里的下标是i,而不是第i号点,这是关键

first_dfs[i]表示第i号点第一次出现在欧拉序列中的位置。

我认为这里主要要理解的是这三个数组的下标,以及存储元素的意义,这里再强调一遍

depth[i]和Euler[i]中的i表示这个点在dfs中是第i个访问到的

比如dfs第4访问到的点是1,1号点的深度为3,则Euler[4]=1,depth[4]=3

而first_dfs[i]中的i表示的就是第i号点,存储的是Euler[k]的下标k,用上面的例子,如果点1第一次访问到的位置是2,则first[1]=2,Euler[2]=1,depth[2]=3

理解了这个,上面的代码就好办了,另外first_dfs还能顺便充当vis数组,因为first_dfs没更新过的点一定没有访问过。

for(int i=1;i<=now_ord;i++){
    st_dep[0][i]=depth[i],st_ord[0][i]=dfs_ord[i];
}
for(int i=1;i<=log(now_ord)/log(2);i++){//cmath自带的log是常用对数,需要套换底公式换到以2为底
    for(int j=1;j<=now_ord-(1<<i);j++){
        if(st_dep[i-1][j]<st_dep[i-1][j+(1<<(i-1))]){
            st_dep[i][j]=st_dep[i-1][j];
            st_ord[i][j]=st_ord[i-1][j]         
         }
        else {
            st_dep[i][j]=st_dep[i-1][j+(1<<(i-1))];
            st_ord[i][j]=st_ord[i-1][j+(1<<(i-1))];         
        }
    }
}

然后是st表部分,这里要对两个st表分别处理.st_dep是深度的st表,st_ord是编号的st表。因为我们在处理st表的时候,要以深度作为判据,但是我们实际要求的并不是深度,而是编号,所以我们再定义一个st表,表示当前具有深度最大值的点是几号点。所以我们可以看到,实际参与比较的是深度,但是ord也要跟着变。st表的处理类似于递推,j到j+2^i的最大值就等于j到j+2^{i-1}的最大值与j+2^{i-1}到j+2^{i-1}+2^{i-1}的最大值中较大的那个。编号也跟着变。这样就处理出了整个st表。

最后是询问的处理

while(m--){
    int a,b;
    scanf("%d%d",&a,&b);
    a=first_dfs[a],b=first_dfs[b];
    //cout<<"first:"<<a<<" first:"<<b<<endl;
    if(a>b)swap(a,b);
    long long k=0;
    //cout<<"Yes"<<endl;
    while((long long)(1<<k)<=b-a+1)k++;
    k--;
    if(st_dep[k][a]<st_dep[k][b-(1<<k)+1]){
        printf("%d\n",st_ord[k][a]);
    }       
    else {
        printf("%d\n",st_ord[k][b-(1<<k)+1]);
    }       
}

首先利用first_dfs找出a与b第一次出现的位置,然后确保a在b前面。

然后while处理出不大于a-b的2的最大幂,但是注意要自减1,因为是判断完了再加,当不满足新一轮判断条件,k已经在旧一轮加上了,所以k不满足条件,会导致超出a-b的边界,所以k--是十分必要的,下面的判断参见上面st表部分。还有一点就是while循环的时候判断条件必须是b-a+1,因为从a到b是包括a和b的闭区间,所以b-a的长度并不是[a-b]的长度,while下面的if判断也是同理。

下面是完整代码

#include<cmath>
#include<cstdio>
#include<cstring>
#include<iostream>
#include<algorithm>
//#define r 1<<(i-1)
using namespace std;
//const int MAXN=500010;
int n,m,s,now_ord=0;

int first[1000010],nextt[1000010],to[1000010];
int first_dfs[1000010],dfs_ord[1000010],depth[1000010];
int st_dep[40][1000100],st_ord[40][1000100];
void add_edge(int start,int end){
    static int ccount=0;
    ccount++;
    nextt[ccount]=first[start];
    first[start]=ccount;
    to[ccount]=end;
}
void dfs(int nowp,int dep){
    //cout<<nowp<<" ";
    now_ord++;
    first_dfs[nowp]=now_ord;
    dfs_ord[now_ord]=nowp;
    depth[now_ord]=dep;
    for(int i=first[nowp];i;i=nextt[i]){
        int v=to[i];
        if(!first_dfs[v]){
            dfs(v,dep+1);
            //cout<<nowp<<" ";
            now_ord++;
            dfs_ord[now_ord]=nowp,depth[now_ord]=dep;
        }
    }
}
int main(){
    cin>>n>>m>>s;
    for(int i=1;i<n;i++){
        int a,b;
        scanf("%d%d",&a,&b);
        add_edge(a,b);
        add_edge(b,a);
    }
    dfs(s,1);
    /*cout<<endl;
    for(int i=1;i<=n;i++){
        cout<<first_dfs[i]<<" ";
    }
    cout<<endl;
    for(int i=1;i<=now_ord;i++){
        cout<<dfs_ord[i]<<" "; 
    }
    cout<<endl;
    for(int i=1;i<=now_ord;i++){
        cout<<depth[i]<<" ";
    }
    cout<<endl;
    */
    for(int i=1;i<=now_ord;i++){
        st_dep[0][i]=depth[i],st_ord[0][i]=dfs_ord[i];
    }
    for(int i=1;i<=log(now_ord)/log(2);i++){
        for(int j=1;j<=now_ord-(1<<i)+1;j++){
            if(st_dep[i-1][j]<st_dep[i-1][j+(1<<(i-1))]){
                st_dep[i][j]=st_dep[i-1][j];
                st_ord[i][j]=st_ord[i-1][j];
            }
            else {
                st_dep[i][j]=st_dep[i-1][j+(1<<(i-1))];
                st_ord[i][j]=st_ord[i-1][j+(1<<(i-1))];
            }
        }
    }
    while(m--){
        int a,b;
        scanf("%d%d",&a,&b);
        a=first_dfs[a],b=first_dfs[b];
        //cout<<"first:"<<a<<" first:"<<b<<endl;
        if(a>b)swap(a,b);
        long long k=0;
        //cout<<"Yes"<<endl;
        while((long long)(1<<k)<=b-a+1)k++;
        k--;
        if(st_dep[k][a]<st_dep[k][b-(1<<k)+1]){
            printf("%d\n",st_ord[k][a]);
        }
        else {
            printf("%d\n",st_ord[k][b-(1<<k)+1]);
        }

    }
    return 0;
}