[[toc]]
学习目标
- 理解 ST 表的原理
- 倍增思想
- ST 表的构建
- ST 表的查询
问题引用
问题描述
有一个长度为
输入格式
第一行两个整数
第二行
接下来的
输出格式
输出
样例输入
5 3
5 2 3 1 6
1 2
2 3
2 5
样例输出
5
3
6
暴力解法
显然可以想到暴力解法,对于每个询问,枚举区间
点击
#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;
}
显然这个代码只能拿到部分分。对于
倍增思想
第一步: 创建ST表
ST表,
它基于,下面的一种基本区间思想
- 一个长度为
的区间 可以平均分成左右长度为 两部分 - 区间最值
可以由 和 得到:
如果我们把区间
这显然是一种DP,因为符合DP的无后效果性原则(DAG).

于是我们想到对于整个序列
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解法
我已经知道任意binary jump的方式来快速得到任意区间的区值.
我们发现了区间信息与跳跃的不一致.比如我想知道区间binary jump的方式是跳,如下:

这样的话,我们不能利用binary jump每条边只走一次,那么每次跳跃正好可以统计个边数为

于是我们按下面的操作把++点转换成边++
点转成边:
- 原数据有
个点 ,新数据有 个点: - 原点
的值 为新的数据上边 的值.
于是我们写出如下的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;
}
得到
这里给三种方案
1. 使用换底公式
int log2(int n) {
return log(n)/log(2);
}
2. 使用dp思想
这里使用数学思想,对于整数
可以想到这个数字
于是我们得到公式
比公式(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) );
}
查询
初始化后,应该如何查询区间最值呢?
重要证明
对于任意区间
这表明,对于任意区间

设长度
根据二进制原理,这是显然的.因为
动画
ST 表的时间复杂度分析
- 初始化每个元素
开头最多做 次,总时间 .其实看一下 数组大小就知道了。 - 显然查询是
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 统一所有的函数的名字