004 - Sum of Max of Sum and Product 解説
実行時間制限: 2sec / メモリ制限: 1024 MB
解説
一般に、正整数 $x,y$ に対して $\max(x + y, xy) = xy$ となる必要十分条件は $x,y$ が共に $1$ でないことです。 逆に、$x,y$ のどちらかが $1$ であるとき、$\max(x + y, xy) = xy + 1$ が成り立ちます。このことから、求める和は $$\sum_{i = 1}^{N} \sum_{j = 1}^N \max(A_i + B_j, A_iB_j) = \sum_{i = 1}^N \sum_{j = 1}^N A_i B_j + (A_i,B_j \text{ の少なくとも一方が } 1\text{ であるような組の個数 } )$$ と表現できます。 前者については $$\sum_{i = 1}^N \sum_{j = 1}^N A_iB_j = \left(\sum_{i = 1}^N A_i\right)\left(\sum_{j = 1}^N B_j\right)$$ と因数分解できるので、$O(N)$ の計算量で求められます。後者については包除原理によって計算できます。 具体的には、数列 $A$ に含まれる $1$ の個数を $x$ 、$B$ に含まれる $1$ の個数を $y$ として、 $ xN + yN - xy$ で表されます。$x,y$ も $O(N)$ で計算できるので、この問題が $O(N)$ で解けました。
コード例
以下はC++による想定解法のコードです。
#include <iostream>
using namespace std;
using ll = long long;
const ll MOD = 998244353;
int main() {
int N;
cin >> N;
ll sum_A = 0, sum_B = 0, x = 0, y = 0;
for (int i = 0; i < N; i++) {
ll A_i;
cin >> A_i;
sum_A = (sum_A + A_i) % MOD;
if (A_i == 1) x++;
}
for (int j = 0; j < N; j++) {
ll B_j;
cin >> B_j;
sum_B = (sum_B + B_j) % MOD;
if (B_j == 1) y++;
}
ll ans = 0;
ans = (sum_A * sum_B + x * N + y * N - x * y) % MOD;
cout << ans << endl;
return 0;
}