ホーム
/
作ったもの
/
考えたこと
/
きろく

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;
}