[[toc]]

学习目标

  • 理解 ST 表的原理
  • 倍增思想
  • ST 表的构建
  • ST 表的查询

问题引用

问题描述

有一个长度为n(n105)n( n \leqslant 10^5)的整数序列,a1,a2,,ana_1,a_2,\cdots,a_n,现在有Q(Q105)Q( Q \leqslant 10^5)个查询,每一个查询给出一个区间[l,r][l,r],求出区间[l,r][l,r]中最大值.

输入格式

第一行两个整数n,Qn,Q,表示序列长度和查询个数.

第二行nn个整数,表示序列a1,a2,,ana_1,a_2,\cdots,a_n.

接下来的QQ行,每行两个整数l,rl,r,表示查询的区间[l,r][l,r].

输出格式

输出QQ行,每行一个字符串,表示查询的区间[l,r][l,r]中最大值.

样例输入

5 3
5 2 3 1 6
1 2
2 3
2 5

样例输出

5
3
6

暴力解法

显然可以想到暴力解法,对于每个询问,枚举区间 [l,r][l,r] 中的最大值,每一次查询的时间复杂度为 O(n)O(n),总时间复杂度为 O(n2)O(n^2)

点击
#include <iostream>
using namespace std;
const int maxn = 1e5 + 5;

int n, q;
int a[maxn];

int query(int l,int r) {
    int t = a[l];
    for (int i = l + 1;i<= r; i++)
    {
        if( t < a[i])
            t = a[i];
    }
    return t;
}

int main() {
    cin >> n >> q;
    for(int i = 1; i <= n; i++)
        cin >> a[i];
    for (int i = 1; i <= q;i++) {
        int l, r;
        cin >> l >> r;
        cout << query(l, r) << "\n";
    }
        return 0;
}

显然这个代码只能拿到部分分。对于10510^5的数据规模会TLE\mathcal{TLE}

倍增思想

第一步: 创建ST表

ST表,spare tablespare\ table本意稀疏矩阵

它基于,下面的一种基本区间思想

  1. 一个长度为2k(k1)2^k(k\geqslant 1)的区间AA可以平均分成左右长度为2k12^{k-1}两部分B,CB,C
  2. 区间最值max(A)max(A)可以由max(B)max(B)max(C)max(C)得到: max(A)=max{max(B),max(C)}max(A)=max\{max(B),max(C)\}

如果我们把区间AA按这种思想不停的分割,最后会得到一个树状结构,每个节点代表一个区间,树的高度为k+1k+1,叶子结点代表长度为202^0,表示原区间上的单个点的值.

这显然是一种DP,因为符合DP的无后效果性原则(DAG).

figure1

于是我们想到对于整个序列[1,n][1,n],我们可以得到这个序列信息f(i,j)f(i,j),表示从位置ii开始的长度为2j2^j的区间的最大值.

TODO: 这里是不是需要一个图,或一个P5的动画来解释一下这个DP运行的过程?

代码如下:

//这里是得到sparse table 的dp算法

const int maxn = 1e5+5;
const int maxk = 50; //这里为什么取50,因为足够大

int f[maxn][maxk]; //定义dp数组

int n; //区间总长度
int a[maxn]; //原始数组

void init_st() {
    //边界
    for(int i=1; i<=n; i++) f[i][0] = a[i];

    // 枚举长度
    for (int k = 1; k <= maxk; k++)
    {
        //枚举起点
        // 这里i+(1<<k) 表示 以i为起点的长度为2^k的区间 的右端点
        // [------ ][there]
        // 这里这样写是为了防止越界
        for (int i = 1; i+(1<<k) <= n+1 ;i++) {
            f[i][k] = max(f[i][k - 1], f[i + 1 << (k - 1)][k - 1]);
        }
    }
}

TODO: 写一个vue的动画

binary jump解法

我已经知道任意f(i,j)f(i,j),也就是长度为2j2^j的区间的最大值信息了.那如何使用binary jump的方式来快速得到任意区间的区值.

我们发现了区间信息与跳跃的不一致.比如我想知道区间[1,4][1,4]的最值,如果采用binary jump的方式是跳,如下:

pos121pos320pos4 pos 1 \xrightarrow{2^1} pos 3 \xrightarrow{2^0} pos4

jump error

这样的话,我们不能利用f(i,j)f(i,j)来统计信息了.本质是jump(i,1)jump(i,1)本质是指从位置ii,走两了步,执行了两次jumpjump函数映射,相当于走了两条边.且binary jump每条边只走一次,那么每次跳跃正好可以统计个边数为2k2^k条边的信息了!

jump right

于是我们按下面的操作把++点转换成边++

点转成边:

  1. 原数据有nn个点1,2,,n1,2,\cdots,n,新数据有n+1n+1个点:0,1,2,,n0,1,2,\cdots,n
  2. 原点ii的值xx为新的数据上边i1,i\langle i-1,i\rangle的值.

于是我们写出如下的binary jump代码

点击
#include <iostream>
using namespace std;
const int maxn = 1e5+5;
const int maxk = 50; //这里为什么取50,因为足够大

int f[maxn][maxk]; //定义dp数组

int n; //区间总长度
int q; // q个query
int a[maxn]; //原始数组

void init_st() {
    //边界
    for(int i=1; i<=n; i++) f[i][0] = a[i];

    // 枚举长度
    for (int k = 1; k <= maxk; k++)
    {
        //枚举起点
        // 这里i+(1<<k) 表示 以i为起点的长度为2^k的区间 的右端点
        // [------ ][there]
        // 这里这样写是为了防止越界
        for (int i = 1; i+(1<<k) <= n+1 ;i++) {
            f[i][k] = max(f[i][k - 1], f[i + 1 << (k - 1)][k - 1]);
        }
    }
}

template<typename T>
T highbit(int n) {
    return sizeof(T) * 8 - 1 - __builtin_clz(n);
}

// 在新数据上位置p是否可达
bool can_reach(int pos,int end) {
    if( pos > n) return  0;
    return pos <= end;
}

//新位置上的点x到跳转2^j
int max_val(int x,int j) {
    return f[x+1][j];
}

//查询 区间的最值,进行binary jump
int query(int l,int r) {
    int start = l - 1; //在新数据上的起点
    int maxk = highbit(n);

    int ans  = a[l]; //初始化ans,为第一个值
    for(int k = maxk;k>=0;k--) {
        int new_pos = start + (1<<k);
        if( can_reach(new_pos,r) )
        {
            ans = max(ans,max_val(start,k));
            start += (1<<k);
        }
    }
    return ans;
}


int main(int argc, char const *argv[])
{
    //读取数据
    cin >> n >> q;
    for(int i =1;i<=n;i++) {
        cin >> a[i];
    }
    init_st();
    for(int i =1;i<=q;i++) {
        int l,r;
        cin >>l>> r;
        cout << query(l,r);
    }
    return 0;
}

得到log2n\lfloor log_2^n \rfloor

这里给三种方案

1. 使用换底公式

int log2(int n) {
    return log(n)/log(2);
}

2. 使用dp思想

这里使用数学思想,对于整数nn来说,它的log2n\lfloor log_2^n \rfloor等价于bb对应的二进制bin(n)bin(n)只保留最高位置的1,其余位置全部置为0后得到的数.

可以想到这个数字highbit(n)highbit(n)对应的为: bin(n)1bin(n) \gg 1后再highbit(n1)+1highbit( n \gg 1) +1.

于是我们得到公式

highbit(n)=highbit(n1)+1 highbit(n) = highbit(n\gg 1) +1

比公式(DP方程)的正确性,可以使用数学归纳法证明.

// 存储每个数字对应的最高位是第几位
int  highbit[maxn]; 
void init_highbit() {
    highbit[0] = -1;
    for(int i = 1; i <= n; i++) {
        highbit[i] = highbit[i>>1] + 1;
    }
}

//得到对应的$\lfloor log_2^n \rfloor$
int log2(int n) {
    return 1 << highbit[n];
}

3. 使用 __builtin_clz

__builtin_clz()是gcc内置的函数,可以得到数字对应的二进制左边的0的数量.

template<typename T
T log2(T n) {
    return 1<< ( sizeof(T) *8 - 1 - __builtin_clz(n) );
}

查询

初始化后,应该如何查询区间最值呢?

重要证明

对于任意区间[l,r][l,r],都有2×log2rl+1rl+12 \times \lfloor log_2^{r-l+1} \rfloor \geqslant r-l+1.

这表明,对于任意区间[l,r][l,r],按如下操作,都能覆盖整个区间.

st_figure2

设长度len=rl+1len = r-l+1,

根据二进制原理,这是显然的.因为log2len\lfloor log_2^{len} \rfloor把最高位的11保留,其它位置全部置为00,这样两个数加起来一定超过lenlen

100100 \def\arraystretch{1.5} \begin{array}{cccc} \boxed{1} & \boxed{0} & \cdots & \boxed{0} \\ \boxed{1} & \boxed{0} & \cdots & \boxed{0} \\ \end{array}

动画

ST 表的时间复杂度分析

  • 初始化每个元素ii开头最多做log2nlog_2^n次,总时间nlognnlogn.其实看一下f(i,j)f(i,j)数组大小就知道了。
  • 显然查询是O(1)O(1)
template<typename T
T log2(T n) {
    return  sizeof(T) *8 - 1 - __builtin_clz(n) ;
}

int query(int l,int r) {
    int k = log2(r-l+1);
    return max(f[i][k],f[i+(1<<k)][k]);
}

ST 表的应用

  • 静态区间最值查询
  • dp优化
  • 其它?TODO

模板代码

综合上面的代码

#include <iostream>
#include <algorithm>
using namespace std;
const int maxn = 1e5+5;
const int maxk = 30;

int n,q;// q表示查询的数量
int a[maxn];
int f[maxn][maxk];

template<typename T>
T log2(int n) {
    return sizeof(T) * 8 - 1 - __builtin_clz(n);
}

//初始化 f 数组
void init_st() {
    for(int i =1 ;i<= n;i++)
        f[i][0] = a[i]
    int LOG = log2(n);
    for(int k = 1; k<=LOG;k++)
    {
        //枚举起点
        for(int i =1; i+(1<<LOG)-1 <=n;i++ )
        {
            f[i][k] = max(f[i][k-1],f[i+(1<<(k-1))][k-1]);
        }
    }
}

int query(int l,int r) {
    int k = log2(r-l+1);
    return max(f[l][k],f[r-(1<<k)+1][k]);
}



int main() {
    cin >> n>>q;
    for(int i =1 ;i<= n;i++) 
        cin >> a[i];
    init_st();

    for(int i =1;i<=q;i++) {
        int l,r;
        cin >> l >> r;
        cout << query(l,r) << endl;
    }
    return 0;
}

TODO 统一所有的函数的名字