Skip to content

线段树主要可以用来优化不断更新的区间求某些统计信息(最大值,最小值等)这一个过程

剑指 Offer 42. 连续子数组的最大和

题目描述

输入一个整型数组,数组中的一个或连续多个整数组成一个子数组。求所有子数组的和的最大值。

要求时间复杂度为O(n)。

示例1:

bash
输入: nums = [-2,1,-3,4,-1,2,1,-5,4]
输出: 6
解释: 连续子数组 [4,-1,2,1] 的和最大,为 6。

思路和算法

这道题除了可以使用动态规划求解外,我们也可以使用分治方法。不仅可以解决区间 \([0, n-1]\),还可以用于解决任意的子区间 \([l,r]\)的问题。如果我们把 \([0, n-1]\)分治下去出现的所有子区间的信息都用堆式存储的方式记忆化下来,即建成一颗真正的树之后,我们就可以在 \(O(\log n)\)的时间内求到任意区间内的答案,我们甚至可以修改序列中的值,做一些简单的维护,之后仍然可以在 \(O(\log n)\)的时间内求到任意区间内的答案,对于大规模查询的情况下,这种方法的优势便体现了出来。这棵树就是一种神奇的数据结构——线段树。

我们定义一个操作 get(a, l, r) 表示查询 a 序列 \([l,r]\)区间内的最大子段和,那么最终我们要求的答案就是 get(nums, 0, nums.size() - 1)。如何分治实现这个操作呢?对于一个区间 \([l,r],\)我们取 \(m = \lfloor \frac{l + r}{2} \rfloor\),对区间 \([l,m]\)\([m+1,r]\)分治求解。当递归逐层深入直到区间长度缩小为 1 的时候,递归「开始回升」。这个时候我们考虑如何通过 \([l,m]\)区间的信息和 \([m+1,r]\)区间的信息合并成区间 \([l,r]\)的信息。

最关键的两个问题是:

  1. 我们要维护区间的哪些信息呢?
  2. 我们如何合并这些信息呢?

对于一个区间\([l,r]\),我们可以维护四个量:

  • lSum 表示 \([l,r]\) 内以 l_l_ 为左端点的最大子段和
  • rSum 表示\( [l,r]\) 内以 r_r_ 为右端点的最大子段和
  • mSum 表示 \([l,r]\)内的最大子段和
  • iSum 表示 \([l,r]\) 的区间和

以下简称 \([l,m]\)\([l,r]\) 的「左子区间」,\([m+1,r]\)\([l,r]\)的「右子区间」。我们考虑如何维护这些量呢(如何通过左右子区间的信息合并得到 \([l,r]\) 的信息)?

对于长度为 1 的区间 \([i, i]\),四个量的值都和 \(\textit{nums}[i\)相等。对于长度大于 1 的区间:

  • 首先最好维护的是 \(\textit{iSum}\),区间 \([l,r]\)\(\textit{iSum}\) 就等于「左子区间」的 \(\textit{iSum}\) 加上「右子区间」的 \(\textit{iSum}\)
  • 对于 \([l,r]\)\(\textit{lSum}\),存在两种可能,它要么等于「左子区间」的 \(\textit{lSum}\),要么等于「左子区间」的 \(\textit{iSum}\) 加上「右子区间」的 \(\textit{lSum}\),二者取大。
  • 对于 \([l,r]\)\(\textit{rSum}\),同理,它要么等于「右子区间」的 \(\textit{rSum}\),要么等于「右子区间」的 \(\textit{iSum}\) 加上「左子区间」的 \(\textit{rSum}\),二者取大。
  • 当计算好上面的三个量之后,就很好计算 \([l,r]\)\(\textit{mSum}\) 了。我们可以考虑 \([l,r]\)\(\textit{mSum}\) 对应的区间是否跨越 m——它可能不跨越 m,也就是说 \([l,r]\)\(\textit{mSum}\) 可能是「左子区间」的 \(\textit{mSum}\) 和 「右子区间」的 \(\textit{mSum}\) 中的一个;它也可能跨越 m,可能是「左子区间」的 \(\textit{rSum}\)和 「右子区间」的 \(\textit{lSum}\) 求和。三者取大。

代码

c
class Solution {
public:
    struct Status {
        int lSum, rSum, mSum, iSum;
    };

    Status pushUp(Status l, Status r) {
        int iSum = l.iSum + r.iSum;
        int lSum = max(l.lSum, l.iSum + r.lSum);
        int rSum = max(r.rSum, r.iSum + l.rSum);
        int mSum = max(max(l.mSum, r.mSum), l.rSum + r.lSum);
        return (Status) {lSum, rSum, mSum, iSum};
    };

    Status get(vector<int> &a, int l, int r) {
        if (l == r) {
            return (Status) {a[l], a[l], a[l], a[l]};
        }
        int m = (l + r) >> 1;
        Status lSub = get(a, l, m);
        Status rSub = get(a, m + 1, r);
        return pushUp(lSub, rSub);
    }

    int maxSubArray(vector<int>& nums) {
        return get(nums, 0, nums.size() - 1).mSum;
    }
};

复杂度分析

假设序列 a 的长度为 n。

时间复杂度:假设我们把递归的过程看作是一颗二叉树的先序遍历,那么这颗二叉树的深度的渐进上界为 \(O(\log n)\),这里的总时间相当于遍历这颗二叉树的所有节点,故总时间的渐进上界是 \(O(\sum_{i=1}^{\log n} 2^{i-1})=O(n)\),故渐进时间复杂度为 \(O(n)\)

空间复杂度:递归会使用 \(O(\log n)\)的栈空间,故渐进空间复杂度为 \(O(\log n)\)

最长递增子序列 II

题目描述

给你一个整数数组 nums 和一个整数 k

找到 nums 中满足以下要求的最长子序列:

  • 子序列 严格递增
  • 子序列中相邻元素的差值 不超过 k

请你返回满足上述要求的 最长子序列 的长度。

子序列 是从一个数组中删除部分元素后,剩余元素不改变顺序得到的数组。

示例 1:

bash
输入:nums = [4,2,1,4,3,4,5,8,15], k = 3
输出:5
解释:
满足要求的最长子序列是 [1,3,4,5,8] 。
子序列长度为 5 ,所以我们返回 5
注意子序列 [1,3,4,5,8,15] 不满足要求,因为 15 - 8 = 7 大于 3 。

示例 2:

bash
输入:nums = [7,4,5,1,8,12,4,7], k = 5
输出:4
解释:
满足要求的最长子序列是 [4,5,8,12] 。
子序列长度为 4 ,所以我们返回 4

示例 3:

bash
输入:nums = [1,5], k = 1
输出:1
解释:
满足要求的最长子序列是 [1] 。
子序列长度为 1 ,所以我们返回 1

提示:

  • 1 <= nums.length <= 10^5
  • 1 <= nums[i], k <= 10^5

思路和算法

在求解「上升子序列」问题时,一般有两种优化方法:

  1. 单调栈 + 二分优化;
  2. 线段树、平衡树等数据结构优化。

这两种做法都可以用 \(O(n\log n)\)的时间解决 300. 最长递增子序列


对于本题,由于有一个差值不超过 \(k\)的约束,用线段树更好处理。

具体来说,定义 \(f[i][j]\)表示 \(\textit{nums}\)\(i\)个元素中以元素\(j\)结尾的满足题目两个条件的子序列的最长长度。

  • \(j\ne\textit{nums}[i]\)时,\(f[i][j] = f[i-1][j]\)
  • \(j=\textit{nums}[i]\) 时,我们可以从 \(f[i-1][j']\)转移过来,这里 \(j-k\le j'<j,\)取最大值,得

\(f[i][j]= 1+ \max_ {j′=j−k}^{j−1}f[i−1][j']\)

上式有一个「区间求最大值」的过程,这非常适合用线段树计算,且由于 \(f[i]\) 只会从 \(f[i-1]\)转移过来,我们可以\(f\)的第一个维度优化掉。这样我们可以用线段树表示整个 \(f\) 数组,在上面查询和更新。

最后答案为 \(\max(f[n-1])\),对应到线段树上就是根节点的值。

注意事项

  1. 这里的线段树划分区间的依据是nums数组中的元素值,按值的范围划分!
  2. 由于nums[i]最小为1,所以线段的范围应该为[1,max(nums[])]
  3. 用数组(堆式存储)来实现线段树时候,数组大小为4*max(nums[])
  4. max[1]代表根节点,假设当前节点下标为o,左子树o*2,右子树o*2+1
  5. 线段树每个节点的线段范围不存储。比如max[1]只维护了每个线段范围内的统计信息,不会去存储区间范围[1,max(nums[])]
  6. 线段树统计信息:假设线段树范围为[l,r],那么max[k]表示每轮迭代的时候 \(\max_ {j′=l}^{r}f[j']\)
  7. base case: nums[i]=1时候,f[1] = 1(对于\({\forall}i\)),因为以1结尾的最长递增子序列只能为[1],其长度为1。

代码:

c
class Solution {
    vector<int> max; 

    // l:线段树区间范围最小值
    // r:线段树区间范围最大值
    // j: 元素值j
    // val:线段树更新维护的信息,这里就是更新后的f[j]
    // o: 线段树数组下标,左子树o*2,右子树o*2+1
    // 修改信息:取区间[l,r]范围内,f[l]~f[r]的最大值
    void modify(int o, int l, int r, int j, int val) {
        if (l == r) {
            max[o] = val;
            return;
        }
        int m = (l + r) / 2;
        if (j <= m) modify(o * 2, l, m, j, val);
        else modify(o * 2 + 1, m + 1, r, j, val);
        max[o] = std::max(max[o * 2], max[o * 2 + 1]);
    }

    // l:线段树范围最小值
    // r:线段树范围最大值
    // o: 线段树数组下标,左子树o*2,右子树o*2+1
    // 返回待查找区间 [L,R] 内的最大值
    int query(int o, int l, int r, int L, int R) { // L 和 R 在整个递归过程中均不变,将其大写,视作常量
        if (L <= l && r <= R) return max[o];
        int res = 0;
        int m = (l + r) / 2;
        if (L <= m) res = query(o * 2, l, m, L, R);
        if (R > m) res = std::max(res, query(o * 2 + 1, m + 1, r, L, R));
        return res;
    }

public:
    int lengthOfLIS(vector<int> &nums, int k) {
        int u = *max_element(nums.begin(), nums.end());
        max.resize(u * 4); //数组大小
        for (int x: nums) {
            if (x == 1) modify(1, 1, u, x, 1);
            else {
                int res = 1 + query(1, 1, u, std::max(x - k, 1), x - 1);
                modify(1, 1, u, x, res);
            }
        }
        return max[1];
    }
};

复杂度分析

  • 时间复杂度:\(O(n\log U)\),其中 \(n\)\(\textit{nums}\) 的长度,\(U=\max(\textit{nums})\)
  • 空间复杂度:\(O(U)\)

用心记录,持续成长