Editorial for Expected Cells (Stretch)
Approach
Do not enumerate pairs of sampled cells. For each grid cell (r,c) let
Ir,c=[(r,c) lies in the bounding rectangle].
The area is A=∑r,cIr,c, so E[A]=∑r,cPr(Ir,c=1).
Cell (r,c) is outside the sampled bounding rectangle if and only if both sampled cells lie strictly to one side of it: both left of column c, both right, both above row r, or both below. Let L,R,U,D be the total weights of those four strict sides, and let NW,NE,SW,SE be the total weights of the four strict quadrants. Writing W for the total weight, inclusion-exclusion gives
~\Pr((r, c)\text{ inside}) = 1
- \frac{L^2 + R^2 + U^2 + D^2}{W^2}
- \frac{NW^2 + NE^2 + SW^2 + SE^2}{W^2}~.
All eight region weights are O(1) from a two-dimensional prefix sum together with the global row and column totals. Summing over every cell is O(nm).
A pairwise enumeration of sampled cells is O((nm)2) and is too slow.
Solution (C++)
#include <bits/stdc++.h>
using namespace std;
const int MOD = 1000000007;
long long modpow(long long a, long long e) {
long long r = 1;
a %= MOD;
while (e) {
if (e & 1) {
r = r * a % MOD;
}
a = a * a % MOD;
e >>= 1;
}
return r;
}
long long sqmod(long long x) {
x %= MOD;
if (x < 0) {
x += MOD;
}
return x * x % MOD;
}
int main() {
ios::sync_with_stdio(false);
cin.tie(0);
int n, m;
cin >> n >> m;
vector<vector<long long> > pref(n + 1, vector<long long>(m + 1, 0));
for (int r = 1; r <= n; r++) {
for (int c = 1; c <= m; c++) {
long long w;
cin >> w;
pref[r][c] = w + pref[r - 1][c] + pref[r][c - 1] - pref[r - 1][c - 1];
}
}
auto rect = [&](int r1, int c1, int r2, int c2) -> long long {
if (r1 > r2 || c1 > c2 || r1 < 1 || c1 < 1 || r2 > n || c2 > m) {
return 0;
}
return pref[r2][c2] - pref[r1 - 1][c2] - pref[r2][c1 - 1] + pref[r1 - 1][c1 - 1];
};
long long W = pref[n][m];
long long W2 = sqmod(W);
long long num = 0;
for (int r = 1; r <= n; r++) {
for (int c = 1; c <= m; c++) {
long long L = rect(1, 1, n, c - 1);
long long R = rect(1, c + 1, n, m);
long long U = rect(1, 1, r - 1, m);
long long D = rect(r + 1, 1, n, m);
long long NW = rect(1, 1, r - 1, c - 1);
long long NE = rect(1, c + 1, r - 1, m);
long long SW = rect(r + 1, 1, n, c - 1);
long long SE = rect(r + 1, c + 1, n, m);
long long bad = (sqmod(L) + sqmod(R) + sqmod(U) + sqmod(D)) % MOD;
long long add = (sqmod(NW) + sqmod(NE) + sqmod(SW) + sqmod(SE)) % MOD;
long long term = (W2 - bad + add) % MOD;
if (term < 0) {
term += MOD;
}
num += term;
if (num >= 4LL * MOD) {
num %= MOD;
}
}
}
num %= MOD;
cout << num * modpow(W2, MOD - 2) % MOD << "\n";
}
Comments0
No comments yet
Be the first to comment.
New comment
Log in to join the discussion.