005 - Sum of Max of Sum and Product 2 解説
実行時間制限: 2sec / メモリ制限: 1024 MB
解説
区間の始点 $i$ を固定し、終点 $j$ を $i$ から $N$ まで動かして計算することを考えます。区間の和 $\sum_{k=i}^j A_k$ は、最大でも $N \times \max(A) = 2 \times 10^5 \times 10^9 = 2 \times 10^{14}$ です。一方で、区間の積 $\prod_{k=i}^j A_k$ は、$A_k \ge 2$ である要素を掛け合わせるたびに少なくとも $2$ 倍になります。したがって、区間内に $2$ 以上の要素が約 $48$ 個($\log_2(2 \times 10^{14}) \approx 47.6$)含まれた時点で、区間の積は確実に $2 \times 10^{14}$ を超え、それ以降は $j$ をどこまで伸ばしても常に積が和を上回ることが保証されます。
この性質を利用し、各 $i$ に対して以下の手順で計算を行います。まず、前計算として「各インデックス以降で最初に現れる $1$ 以外の要素の位置」を配列などに記録しておきます。始点 $i$ を固定したら、この前計算を利用して $1$ を読み飛ばしながら $j$ をジャンプさせて探索を進めます。
探索の過程で $1$ が連続する区間が現れた場合、その区間内では「直前までの積は変化せず、和は $1$ ずつ増加する」という一次関数的な推移になります。そのため、区間内の各終点における $\max(\text{和}, \text{積})$ の合計は、場合分けと等差数列の和の公式等を用いて $O(1)$ でまとめて計算することができます。
$2$ 以上の要素を掛け合わせながら探索を進め、積が $2 \times 10^{14}$ を超えた時点でシミュレーションを打ち切ります。打ち切った以降の残りの終点 $j$ については、全て $\max(\text{和}, \text{積}) = \text{積}$ となるため、単なる積の総和を求めればよいことになります。これは後ろから計算するDPなどで「$k$ を始点とする終端までの積の総和」を $O(N)$ で前計算しておくことで、残りの計算も $O(1)$ で求めることができます(この部分の計算は $\bmod 998244353$ の下で行うことに注意してください)。
以上のアルゴリズムにより、各始点 $i$ に対して最大でも $48$ 回程度のジャンプと $O(1)$ の処理を行うだけで答えが求まります。全体の計算量は $O(N \log(N \max A))$ となり、十分に高速にこの問題を解くことができます。
コード例
以下はC++による想定解法のコードです。
#include <iostream>
#include <vector>
using namespace std;
using ll = long long;
const ll MOD = 998244353;
int main() {
int N;
cin >> N;
vector<ll> A(N);
for (int i = 0; i < N; i++) {
cin >> A[i];
}
vector<ll> S(N + 1, 0);
vector<int> next_non1(N + 1, N);
for (int i = N - 1; i >= 0; i--) {
S[i] = (A[i] + (A[i] * S[i + 1]) % MOD) % MOD;
if (A[i] != 1) next_non1[i] = i;
else next_non1[i] = next_non1[i + 1];
}
ll ans = 0;
ll LIMIT = 200000LL * 1000000000LL; // 和の最大値 (2 * 10^14)
ll INV2 = 499122177; // mod 998244353 における 2 の逆元
for (int i = 0; i < N; i++) {
ll current_sum = 0, current_prod = 1, current_prod_mod = 1;
int curr = i;
while (curr < N) {
int nxt = next_non1[curr];
ll ones = nxt - curr;
if (ones > 0) {
ll kth = current_prod - current_sum;
if (kth <= 0) {
// 全て「和」が勝つ場合
ll first_term = (current_sum + 1) % MOD;
ll last_term = (current_sum + ones) % MOD;
ll sum_seq = (first_term + last_term) % MOD * ones % MOD * INV2 % MOD;
ans = (ans + sum_seq) % MOD;
} else if (kth >= ones) {
// 全て「積」が勝つ場合
ans = (ans + current_prod_mod * ones) % MOD;
} else {
// 途中まで「積」、途中から「和」が勝つ場合
ans = (ans + current_prod_mod * kth) % MOD;
ll count_sum = ones - kth;
ll first_term = (current_sum + kth + 1) % MOD;
ll last_term = (current_sum + ones) % MOD;
ll sum_seq = (first_term + last_term) % MOD * count_sum % MOD * INV2 % MOD;
ans = (ans + sum_seq) % MOD;
}
current_sum += ones;
curr = nxt;
}
if (curr == N) break;
// 次の要素を掛けると 2*10^14 を超えるか判定 (オーバーフロー回避のため割り算で判定)
if (LIMIT / current_prod < A[curr]) {
// 以降は確実に「積」が勝つため、前計算した累積積を使って O(1) で処理して打ち切り
ll add = (current_prod_mod * S[curr]) % MOD;
ans = (ans + add) % MOD;
break;
} else {
// 2以上の要素を掛け合わせて更新
current_sum += A[curr];
current_prod *= A[curr];
current_prod_mod = (current_prod_mod * (A[curr] % MOD)) % MOD;
if (current_prod >= current_sum) ans = (ans + current_prod_mod) % MOD;
else ans = (ans + current_sum % MOD) % MOD;
curr++;
}
}
}
cout << ans << endl;
return 0;
}