689. Maximum Sum of 3 Non-Overlapping Subarrays

📋 Đề Bài

Given an integer array nums and an integer k, find three non-overlapping subarrays of length k with maximum sum and return them.

Return the result as a list of indices representing the starting position of each interval (0-indexed). If there are multiple answers, return the lexicographically smallest one.

 

Example 1:

Input: nums = [1,2,1,2,6,7,5,1], k = 2
Output: [0,3,5]
Explanation: Subarrays [1, 2], [2, 6], [7, 5] correspond to the starting indices [0, 3, 5].
We could have also taken [2, 1], but an answer of [1, 3, 5] would be lexicographically larger.

Example 2:

Input: nums = [1,2,1,2,1,2,1,2,1], k = 2
Output: [0,2,4]

 

Constraints:

  • 1 <= nums.length <= 2 * 104
  • 1 <= nums[i] < 216
  • 1 <= k <= floor(nums.length / 3)

🧠 Thuật Toán & Kỹ Thuật

Binary Search (Tìm kiếm nhị phân)Bit Manipulation (Thao tác bit)
⏱️ Thời gian O(log n)
💾 Không gian O(n)

💻 Lời Giải

C++ 0689-maximum-sum-of-3-non-overlapping-subarrays.cpp
struct SegmentTree {
private:
    vector<int> tree, nums;
    
public:
    SegmentTree() {
        
    }
    
    SegmentTree(int n, vector<int> &nums) {
        tree.resize(4*n, -1);
        this->nums = nums;
        for (int i = 0; i < n; ++i) {
            update(1, 0, n - 1, i);
        }
    }
    
    void update(int node, int left, int right, int index) {
        if (index < left || right < index) {
            return;
        }
        if (left == right) {
            tree[node] = index;
            return;
        }
        int mid = (left + right) >> 1;
        update(node*2, left, mid, index);
        update(node*2 + 1, mid + 1, right, index);
        if (tree[node*2] == -1) {
            tree[node] = tree[node*2 + 1];
        }
        else if (tree[node*2 + 1] == -1) {
            tree[node] = tree[node*2];
        }
        else if (nums[tree[node*2]] >= nums[tree[node*2 + 1]]) {
            tree[node] = tree[node*2];
        }
        else {
            tree[node] = tree[node*2 + 1];
        }
    }
    
    int get(int node, int left, int right, int q_left, int q_right) {
        if (q_left > right || q_right < left) {
            return -1;
        }
        if (q_left <= left && right <= q_right) {
            return tree[node];
        }
        int mid = (left + right) >> 1;
        int node_left = get(node*2, left, mid, q_left, q_right);
        int node_right = get(node*2 + 1, mid + 1, right, q_left, q_right);
        if (node_left == -1) {
            return node_right;
        }
        else if (node_right == -1) {
            return node_left;
        }
        else if (nums[node_left] >= nums[node_right]) {
            return node_left;
        }
        else {
            return node_right;
        }
    }
};

class Solution {
public:
    vector<int> maxSumOfThreeSubarrays(vector<int>& nums, int k) {
        const int n = nums.size();
        int s = 0;
        vector<int> nums1(n, -1), nums2(n, -1);
        
        for (int i = 0, j = 0; i < n; ++i) {
            s += nums[i];
            if (i - j + 1 == k) {
                nums1[i] = s;
                nums2[j] = s;
                s -= nums[j];
                j++;
            }
        }
        
        SegmentTree st_left(n, nums1), st_right(n, nums2);
                
        vector<int> ans;
        int max_val = 0;
        s = 0;
        
        for (int i = 0, j = 0; i < n; ++i) {
            s += nums[i];
            if (i - j + 1 == k) {
                int get_query_left = st_left.get(1, 0, n - 1, 0, max(0, j - 1));
                int get_query_right = st_right.get(1, 0, n - 1, min(n - 1, i + 1), n - 1);
                if (j >= k && i + k < n && get_query_left != -1 && get_query_right != -1 && \
                    nums1[get_query_left] != -1 && nums2[get_query_right] != -1) {
                    int val = s + nums1[get_query_left] + nums2[get_query_right];
                    if (max_val < val) {
                        max_val = val;
                        ans = {get_query_left - k + 1, j, get_query_right};
                    }
                }
                s -= nums[j];
                j++;
            }
        }
        
        return ans;
    }
};