Hi everyone!
Thanks to K1o0n and Noobish_Monk for the round. This problem looks scary at first, but the final idea is short. I will explain it slowly, step by step, so beginners can follow too. If something is unclear, ask in the comments!
What does the problem ask?
We have a table of numbers with n rows and m columns. Pick any rectangle inside it (a "submatrix"). Let its height be a, its width be b, and the sum of its numbers be S.
- If a + b is even, its price is +S2.
- If a + b is odd, its price is -S2.
We need the total price of all rectangles, modulo 1000000007. Also n * m <= 106, so we need roughly an O(n*m) solution.
Why is brute force too slow?
A table has about (n2 * m2) / 4 rectangles. For n * m = 106 that is far too many. We need a smarter idea.
A tiny example first
Take one row with 3 numbers: [a, b, c]. The rectangles are:
- length 1 (even): +a2, +b2, +c2
- length 2 (odd): -(a+b)2, -(b+c)2
- length 3 (even): +(a+b+c)2
Here the height is 1, so the sign depends on the length. If you expand and add everything, almost all terms cancel and you get:
(a + c)2
The middle number b disappeared completely! Most things cancel, and we only need to find what survives.
Step 1: Break S2 into pairs
S2 means "sum of all numbers" times "sum of all numbers". So it equals the sum of A(p) * A(q) over every pair of cells (p, q) inside the rectangle. Both orders (p, q) and (q, p) count, and p = q is allowed.
Now flip the viewpoint:
Instead of "for each rectangle, look at all pairs inside it", do "for each pair of cells, look at all rectangles that contain both".
For each pair we only need the total of the signs of all rectangles containing both cells. Call it W. Then answer = sum of A(p) * A(q) * W over all pairs.
Step 2: Rows and columns are independent
A rectangle is chosen by top row r1, bottom row r2, left column c1, right column c2. The sign is
(-1)r1 + r2 + c1 + c2
because a = r2 — r1 + 1 and b = c2 — c1 + 1, and only the parity of a + b matters. So the sign splits into a row part and a column part, and W = (row part) * (column part). We solve rows alone, then columns alone.
Step 3: Solve the rows (columns are the same)
Let the two cells have rows lo (smaller) and hi (larger). A rectangle contains both cells if r1 <= lo and r2 >= hi. So the row part is
(sum of (-1)r1 for r1 = 1..lo) * (sum of (-1)r2 for r2 = hi..n)
These are sums like -1 +1 -1 +1 ..., where neighbours cancel:
- First part: if lo is even, everything cancels and it is 0. If lo is odd, one -1 is left, so it is -1.
- Second part: if hi and n have different parity, the number of terms is even and it cancels to 0. If they have the same parity, one term is left and it is (-1)n.
Conclusion for rows: a pair contributes only if
- the smaller row is odd, and
- the larger row has the same parity as n.
In that case the row part is -(-1)n. Columns work the same way with m.
Multiplying the row part and the column part gives (-1)n+m. This is a constant, so
answer = (-1)n+m * (sum of A(p) * A(q) over valid pairs)
A pair is valid if: smaller row is odd, larger row has the parity of n, smaller column is odd, larger column has the parity of m.
That is why the code ends with "if n+m is odd, negate the total".
Check with the example [a, b, c] (n = 1, m = 3): valid column pairs are (1,1), (1,3), (3,1), (3,3), which gives a2 + 2ac + c2 = (a+c)2. It matches!
Step 4: Add up the valid pairs fast
We cannot check all pairs (too slow), so we use suffix sums.
For every ordered pair, one cell has the smaller row and one has the smaller column. That gives 4 cases. In each case we know exactly which parity class each cell must be in. We handle ties (same row or same column) carefully so that no pair is counted twice.
In all 4 cases, the cell we loop over has an odd row, and we need:
"the sum of all good cells in the region below it (to the right or to the left)"
A 2D suffix sum gives this in O(1). We build two tables:
- S0: suffix sums going down-right, over cells with row parity = n and column parity = m.
- S1: suffix sums going down-left, over cells with row parity = n and column odd.
Then for every cell in an odd row we multiply its value by the right table lookup for each case and add it to the total. The exact lookups are in the code comments.
Complexity
- Time: O(n*m) per test
- Memory: O(n*m)
The sum of n * m over all tests is at most 106, so it runs comfortably in 1.5 seconds.
Code
Common mistakes
- Forgetting to make negative numbers non-negative before taking modulo (x %= MOD; if (x < 0) x += MOD;).
- Forgetting that after subtracting in the suffix sums the value can become negative, so we add MOD again.
- Using int for products. Use long long.
Key takeaways
- When you see S2, think "sum over pairs of cells".
- Swap the order: count how much each pair contributes instead of each rectangle.
- Alternating signs (-1 +1 -1 ...) cancel in pairs, so most things become 0.
- Use 2D suffix sums to count what remains in O(1) per cell.
Where I got stuck
[Replace this line with one honest sentence, or delete this whole section.]
Thanks for reading! If it helped, an upvote is appreciated, and feel free to point out any mistakes.



