O(n) 中的解决方案
C++ 源代码如下,程序读取标准输入或
文件,其路径作为命令行的第一个参数提供。
输入文件的格式预计为:
- N 一个整数,表示数组中的元素个数
- N 个整数,用空格分隔,表示数组的元素
编译通过:
g++ -std=c++14 -g -Wall -O0 solution.cpp -o solution
程序将首先使用O(n) 算法计算总和,然后使用
O(n^3)算法验证。
示例运行:
$ ./solution.exe
3
4 5 6
O(n) sum: 32
O(n^3) sum: 32
源代码:
#include <cstdio>
#include <algorithm>
#include <iomanip>
#include <iostream>
#include <vector>
#include <stack>
using namespace std;
int main(int argc, char* argv[]) {
if(argc > 1)
freopen(argv[1], "r", stdin);
// load input array
int N;
cin >> N;
vector<int> A(N);
for(auto& ai : A)
cin >> ai;
// compute sum max of all subarrays in O(n)
vector<int> l(N);
vector<int> r(N);
stack<int> s;
for(int i=0; i<N; ++i) {
while(s.size() && A[s.top()] < A[i]) {
r[s.top()] = i;
s.pop();
}
s.push(i);
}
while(s.size()) {
r[s.top()] = N;
s.pop();
}
for(int i=N-1; i>=0; --i) {
while(s.size() && A[s.top()] <= A[i]) {
l[s.top()] = i;
s.pop();
}
s.push(i);
}
while(s.size()) {
l[s.top()] = -1;
s.pop();
}
int sum = 0;
for(int i=0; i<N; ++i) {
int cs = A[i]*(i-l[i])*(r[i]-i);
sum += cs;
}
cout << "O(n) sum: " << sum << '\n';
// compute sum using O(n^3) algorithm for verification
sum = 0;
for(int i=0; i<N; ++i) {
for(int j=i; j<N; ++j) {
int cs = *max_element(begin(A)+i, begin(A)+j+1);
sum += cs;
}
}
cout << "O(n^3) sum: " << sum << '\n';
}
解决方案的证明
首先,这不是完整的证明。我有一个证据,但它太多了
涉及数学符号被包含在没有 mathjs 的网站上
支持...我会给出证明的草图并让细节给
读者(我知道它很蹩脚)。
解决方案使用多个技巧:
- 所有子数组的最大值之和等于每个子数组的乘积之和
最大值由这是最大值的子数组的数量
- 可以将所有子数组的集合划分为一组
子数组。此分区中的每个集合仅包含具有相同
最大元素和集合的大小很容易计算
让我们命名问题的元素:
- 调用初始数组:
A。值 A[i] 是数组中从 0 开始的索引 i 处的值。
- 数组的大小是
n。该数组从索引0 延伸到n-1。
首先,我定义一个子数组的leader,这是索引i
子数组中的一个元素,使得值A[i] 是最大的
子数组和子数组中索引j < i 处的每个元素都有
A[j]<A[i]。直观地说,子数组的 leader 是它的第一个索引
最大值。
我说 leader 的平等定义了一个equivalence relation
子数组。 (证明留作练习)。
由此,我们知道集合的等价类form a partition
所有子数组。此外,等价类中的所有元素都具有
相同的最大值(由于 leader 函数的定义)。
等价类E_i的大小,leader的所有子数组的集合
是i,很容易从值中计算出来:
-
l(i) 这是i 左侧的第一个索引,其中A[l(i)] >= A[i] 或
-1 如果不存在这样的索引
-
r(i) 这是i 右侧的第一个索引,其中A[l(i)] > A[i] 或
n 如果不存在这样的索引
使用这些符号,E_i 的基数是:(i-l(i))*(r(i)-i)。证明留给读者作为练习。
现在是计算值l(i) 和r(i) 的编程技巧。作为
计算几乎相同,我将只解释l(i) 的计算。我们
维护一组索引,具有以下不变量:
- 堆栈中的索引代表领导者
- 所有索引都按升序排列
- 与堆栈中的索引关联的所有值也按升序排列
我们从左到右扫描数组。对于每个索引i,我们检查它的值是否
大于当前栈顶的值。如果是这种情况,则意味着 leader 位于堆栈顶部的子数组不能超出右侧的当前索引。所以我们将 r(top of stack) 的值更新为i。我们弹出堆栈的顶部,因为它可能不涉及任何超过i 的子数组。
我们继续更新 leaders 并弹出堆栈的顶部,直到
堆栈为空或A[top of stack] >= A[i]。然后我们将i 压入堆栈。
当到达数组的末尾时,堆栈中可能仍有一些索引。
这意味着它们参与延伸到阵列末端的子阵列。
我们将他们的r 值更新为N。
整个扫描更新O(n) 中r() 的所有值。这是因为每个
元素是
因为我们不能多次弹出一个元素,所以内部的while 循环不会运行
整个阵列扫描超过n 次。
使用相同的过程来计算l(),除了:
- 我们从右到左向后扫描
- 由于 leader 的定义,我们使用严格的弱比较
- 我们使用
-1来表示limit是数组的边界
然后我们可以应用公式来计算等价的大小
类并在我们的总结中使用它。导致O(n) 算法,因为我们
需要:
-
O(1)阅读r(i)
-
O(1)阅读l(i)
-
O(1) 读取A(i) E_i 中子数组的最大值