ST 表


算法笔记:倍增思想与 ST 表进阶应用

倍增思想的核心在于利用 $2^i$ 的步长进行预处理和跳转,将线性的 $O(N)$ 遍历优化至 $O(\log N)$ 或 $O(1)$。

1. 基础模板:ST表求区间最值 (RMQ)

关联文件模板.cpp, P_2880_USACO_07_JAN_Balanced_Lineup_G.cpp

  • 题面抽象:给定一个静态数组,多次独立询问区间 $[l, r]$ 内的最大值与最小值之差。
  • 核心思路
  • 状态定义Max[i][j] 表示从索引 $i$ 开始,长度为 $2^j$ 的区间内的最大值。
  • 预处理 (DP):利用区间可加性,将长度为 $2^j$ 的区间拆分为两个长度为 $2^{j-1}$ 的区间。
    状态转移方程:Max[i][j] = max(Max[i][j-1], Max[i + (1<<(j-1))][j-1])。预处理复杂度 $O(n \log n)$。
  • $O(1)$ 查询:对于查询区间 $[l, r]$,找到满足 $2^k \le r - l + 1$ 的最大 $k$ 值。分别从 $l$ 向右和从 $r$ 向左取长度为 $2^k$ 的两个区间,求最大值/最小值。这两个区间必定完全覆盖 $[l, r]$,且最值查询允许区间重叠(即 $max(x, x) = x$ 的幂等性)。
点击查看代码
#define lg __lg
//必须要这么写 要不然 n 会越界
struct ST {
    int n, k;
    vector<int> a;
    vvi Max, Min;
    ST(int n) {
        this->n = n;
        this->k = lg(n);
        a.resize(n + 1);
        Max.resize(n + 1, vi(k + 1));
        Min.resize(n + 1, vi(k + 1));
    }
    void init() {
        for (int i = 1; i <= n; i++) {
            Max[i][0] = Min[i][0] = a[i];
        }
     
        for (int j = 1; j <= k; j++) {
            for (int i = 1; i + (1 << j) - 1 <= n; i++) {
                Max[i][j] = max(Max[i][j - 1], Max[i + (1 << (j - 1))][j - 1]);
                Min[i][j] = min(Min[i][j - 1], Min[i + (1 << (j - 1))][j - 1]);
            }
        }
    }
    int get(int l, int r) {
        if (l > r) swap(l, r);
        int k = lg(r - l + 1);
        return max(Max[l][k], Max[r - (1 << k) + 1][k]) - min(Min[l][k], Min[r - (1 << k) + 1][k]);
    }
};


2. 基础扩展:ST表求区间 GCD

关联文件P_1890_gcd_区间.cpp

  • 题面抽象:给定一个静态数组,多次查询区间 $[l, r]$ 内所有元素的最大公约数 (GCD)
  • 核心思路
  • GCD 操作与求最值一样,完全满足结合律幂等性,即 $gcd(x, x) = x$。
  • 因此,ST 表可以原封不动地套用,只需将 max/min 函数替换为 __gcd 即可。
  • 同样实现 $O(n \log n)$ 预处理,$O(1)$ 查询区间 GCD。
点击查看代码
#define lg __lg
#define gcd __gcd
struct ST {
    int n, k;
    vector<int> a;
    vvi G;
    ST(int n) {
        this->n = n;
        this->k = lg(n);
        a.resize(n + 1);
        G.resize(n + 1, vi(k + 1));
        // Min.resize(n + 1, vi(k + 1));
    }
    void init() {
        for (int i = 1; i <= n; i++) {
            G[i][0] = a[i];
        }
        for (int j = 1; j <= k; j++) {
            for (int i = 1; i + (1 << j) - 1 <= n; i++) {
                G[i][j] = gcd(G[i][j - 1], G[i + (1 << (j - 1))][j - 1]);
            }
        }
    }
    int get(int l, int r) {
        if (l > r) swap(l, r);
        int k = lg(r - l + 1);
        return gcd(G[l][k], G[r - (1 << k) + 1][k]);
    }
};
void solve() {
    int n, q;
    cin >> n >> q;

    ST st(n);
    for (int i = 1; i <= n; i++) {
        cin >> st.a[i];
        // cout << st.a[i] << endl;
    }
    st.init();

    while (q--) {
        int l, r;
        cin >> l >> r;
        cout << st.get(l, r) << endl;
    }
}

3. 树上倍增:LCA 与条件祖先跳转

关联文件树.cpp

  • 题面抽象:给定一棵树,节点有权值。支持两种操作:1. 寻找节点 $u$ 的深度最深的且权值 $\ge w$ 的祖先节点。 2. 在满足“节点权值必须小于等于父节点、大于等于所有子节点”的单调性前提下,更新某个节点的权值。
  • 核心思路
  • 树上倍增 (LCA 基础)f[x][i] 记录节点 $x$ 向上走 $2^i$ 步到达的祖先。DFS 预处理时转移:f[x][i] = f[f[x][i-1]][i-1]
  • 倍增跳表 (核心应用):因为题目限制了父节点权值 $\ge$ 子节点权值,所以从子节点向根节点方向走,权值是单调递增的
  • 要找权值 $\ge w$ 的最深祖先,我们可以反向思考:利用倍增数组从大步数向小步数枚举(从 $i = \log(\text{dep})$ 降到 $0$),如果跳到的祖先权值 < w,说明还没跳够,就贪心地跳过去更新当前节点 u = f[u][i]。最终停下来的节点的前驱(父节点 f[u][0])就是刚好满足 $\ge w$ 的最深祖先。这种方法将树上搜索的复杂度降为了 $O(\log n)$。
点击查看代码
#define lg __lg
struct Tree {
    int n;
    vector<vector<int>> g;
    vector<array<int, 21>> f;
    vector<int> dep;
    vector<int> c;
    Tree(int n) {
        this->n = n;
        g.resize(n + 1);
        f.resize(n + 1);
        dep.resize(n + 1);
        c.resize(n + 1);
    }
    void add(int x, int y) {
        g[x].emplace_back(y);
        g[y].emplace_back(x);
    }
    void dfs(int x, int fa) {
        f[x][0] = fa;
        dep[x] = dep[fa] + 1;
        for (int i = 1; i <= lg(dep[x]); i++) {
            f[x][i] = f[f[x][i - 1]][i - 1];
        }
        for (auto y : g[x]) {
            if (y == fa) continue;
            dfs(y, x);
        }
    }
    int lca(int x, int y) {
        if (dep[x] < dep[y]) swap(x, y);
        while (dep[x] > dep[y]) {
            x = f[x][lg(dep[x] - dep[y])];
        }
        if (x == y) return x;
        for (int i = lg(dep[x]); i >= 0; i--) {
            if (f[x][i] == f[y][i]) continue;
            x = f[x][i];
            y = f[y][i];
        }
        return f[x][0];
    }
	//跳表正确用法
    int find_node(int u, int w) {
        if (c[u] >= w) return u;
        if (c[1] < w) return -1;
        for (int i = lg(dep[u]); i >= 0; i--) {
            if (f[u][0] != 0 && c[f[u][i]] < w) u = f[u][i];
        }
        return f[u][0];
    }
    bool update(int x, int v) {
        if (x != 1 && c[x] + v > c[f[x][0]]) {
            return 0;
        }
        if (g[x].size() != 1) {

            int mx = -1e18;
            for (auto v : g[x]) {
                if (v == f[x][0]) continue;
                mx = max(c[v], mx);
            }
            if (c[x] + v < mx) {
                return 0;
            }
        }
        return 1;
    }

    void work(int rt = 1) {
        dfs(rt, 0);
    }
};
void solve() {
    int n, q;
    cin >> n >> q;
    Tree tr(n);
    for (int i = 1; i <= n; i++) {
        cin >> tr.c[i];
    }
    for (int i = 1; i < n; i++) {
        int u, v;
        cin >> u >> v;
        tr.add(u, v);
    }
    tr.work();
    while (q--) {
        int op, a, b;
        cin >> op >> a >> b;
        if (op == 1) {
            cout << tr.find_node(a, b);
        } else {
            if (tr.update(a, b)) {
                tr.c[a] = tr.c[a] + b;
                cout << "SUCCESS";
            } else
                cout << "FAILED";
        }
        cout << endl;
    }
}


4. 综合进阶:ST表 + 二分查找 (利用单调性)

关联文件GD 终极节奏实验室 .cpp

  • 题面抽象:给定一个数组,对于每个元素 $a[i]$,求出有多少个连续子区间是以 $a[i]$ 作为该区间的总体 GCD 的。
  • 核心思路
  • GCD 的单调性:一个区间的长度越长,其包含的元素越多,整个区间的 GCD 必定单调不增。如果区间 GCD 等于 $a[i]$,说明区间内所有元素都是 $a[i]$ 的倍数。
  • 分离左右边界:由于单调性存在,我们可以固定元素 $a[i]$ 的位置 $i$,分别向左和向右使用二分查找
  • 二分 + $O(1)$ 判定:在二分枚举左右边界 mid 时,需要快速求出 $[mid, i]$ 或 $[i, mid]$ 的 GCD 以判断是否等于 $a[i]$。此时预先建好的 ST 表就派上了用场,提供了 $O(1)$ 的判定能力。
  • 答案统计:找到最远的合法左边界 $ll$ 和最远的合法右边界 $rr$ 后,跨越 $i$ 且满足条件的子区间个数即为组合数学乘法原理:$(i - ll + 1) \times (rr - i + 1)$。该算法整体复杂度为 $O(n \log n)$。
点击查看代码
#define lg __lg
#define gcd __gcd
struct ST {
    int n, k;
    vector<int> a;
    vvi G;
    ST(int n) {
        this->n = n;
        this->k = lg(n);
        a.resize(n + 1);
        G.resize(n + 1, vi(k + 1));
        // Min.resize(n + 1, vi(k + 1));
    }
    void init() {
        for (int i = 1; i <= n; i++) {
            G[i][0] = a[i];
        }
        for (int j = 1; j <= k; j++) {
            for (int i = 1; i + (1 << j) - 1 <= n; i++) {
                G[i][j] = gcd(G[i][j - 1], G[i + (1 << (j - 1))][j - 1]);
            }
        }
    }
    int get(int l, int r) {
        if (l > r) swap(l, r);
        int k = lg(r - l + 1);
        return gcd(G[l][k], G[r - (1 << k) + 1][k]);
    }
};
void solve() {
    int n;
    cin >> n;
    vi a(n + 1);
    ST st(n);
    for (int i = 1; i <= n; i++) {
        cin >> a[i];
        st.a[i] = a[i];
    }
    st.init();
    int ans = 0;
    map<int, int> mp;
    for (int i = 1; i <= n; i++) {

        int l = mp[a[i]], r = i + 1;
        mp[a[i]] = i;
        // 左走
        while (l + 1 != r) {
            int mid = (l + r) >> 1;
            if (st.get(mid, i) == a[i]) {
                r = mid;
            } else
                l = mid;
        }

        int ll = r;
        // 右走
        l = i - 1, r = n + 1;
        while (l + 1 != r) {
            int mid = (l + r) >> 1;
            if (st.get(i, mid) == a[i]) {
                l = mid;
            } else
                r = mid;
        }
        int rr = l;
        int sum = (rr - i + 1) * (i - ll + 1);
        ans += sum;
    }
    cout << ans << endl;
}

不用二分的写法

点击查看代码
// 寻找满足区间 GCD 等于 a[i] 的最远左边界 L
int find_farthest_L(int i) {
    // 如果已经在最左边,没得跳了,直接返回 1
    if (i == 1) return 1;

    int pos = i;          // 当前停留的位置
    int cur_gcd = a[i];   // 当前累积的 GCD
    
    // i 左边最多有 i - 1 个元素,所以最大步长指数是 __lg(i - 1)
    for (int k = __lg(i - 1); k >= 0; k--) {
        
        // 这一段 2^k 长度的区间起点
        int next_start = pos - (1 << k);
        
        // 1. 判断往左跳是否会越过数组边界(起点不能小于 1)
        if (next_start >= 1) {
            
            // 2. O(1) 获取这一整段的 GCD
            int next_chunk_gcd = st[next_start][k];
            
            // 3. 验证合并后是否满足条件
            if (__gcd(cur_gcd, next_chunk_gcd) == a[i]) {
                // 满足条件,放心向左跳跃
                cur_gcd = __gcd(cur_gcd, next_chunk_gcd);
                pos -= (1 << k); // 位置向左移动 2^k
            }
        }
    }
    
    return pos; // 返回最远左边界
}
int find_farthest_R(int i, int n) {
    int pos = i;          // 当前停留的位置
    int cur_gcd = a[i];   // 当前累积的 GCD
    
    // 直接使用 __lg() 计算当前剩余长度能跨越的最大 2^k 步长
    // 细节保障:只要 i <= n,(n - i + 1) 就绝对 >= 1,不会触发 __lg(0) 的未定义行为
    for (int k = __lg(n - i + 1); k >= 0; k--) {
        
        int next_start = pos + 1;
        
        // 1. 判断加上 2^k 长度后是否越界
        if (next_start + (1 << k) - 1 <= n) {
            
            // 2. O(1) 获取这一整段的 GCD
            int next_chunk_gcd = st[next_start][k];
            
            // 3. 验证合并后是否满足条件
            if (__gcd(cur_gcd, next_chunk_gcd) == a[i]) {
                // 满足条件,放心跳跃
                cur_gcd = __gcd(cur_gcd, next_chunk_gcd);
                pos += (1 << k); 
            }
        }
    }
    
    return pos; // 返回最远右边界
}

// 假设主循环遍历每一个 i
int ans = 0;
for (int i = 1; i <= n; i++) {
    int L = find_farthest_L(i);     // O(log N)
    int R = find_farthest_R(i, n);  // O(log N)
    
    // 包含位置 i,且左端点在 [L, i],右端点在 [i, R] 的所有区间
    int left_choices = i - L + 1;
    int right_choices = R - i + 1;
    
    ans += left_choices * right_choices;
}
posted @ 2026-05-16 19:32  r_123  阅读(16)  评论(0)    收藏  举报