【题解】ZJOI2009-对称的正方形
鉴于之前写的代码太垃圾,所以美化了格式!
【ZJOI 2009 Day2】对称的正方形
问题是给定 n * m 的矩阵求其中有多少个上下对称和左右对称的正方形子矩阵。
先看一下样例。
5 5
4 2 4 4 4
3 1 4 4 3
3 5 3 3 3
3 1 5 3 3
4 2 1 2 4
显然是 \(25\) 个单独的和 \(2*2\) 的 \(4\) 和 \(3\)。
一共是 \(27\) 个,这是对的。
Sol1
先考虑暴力。
判断 is(x, y, k) 是 \(O(n^2)\) 的。
枚举 \(x, y, k\) 即可。
复杂度 \(O(n^5)\)。
Sol2
尝试优化枚举的复杂度。
假设一个地方为中心点(可能在格子外),最大扩展 \(k\)。
那么的话比 \(k\) 小的所有的正方形肯定是合法的。
所以二分枚举 \(k\) 即可。
时间复杂度 \(O(n^4 \log n)\)。
Sol3
尝试优化判断的复杂度。
看到回文想到 Manacher 算法。
我们可以列出 \(r_{i, j}\) 和 \(c_{i, j}\) 作为横向和竖向的最长延伸的长度。
所以我们现在可以枚举中心点 \(i, j\) 以及二分长度 \(k\):
从 \(r_{i - k \dots i + k}\) 和 \(c_{j - k \dots j + k}\) 中取 \(\min\) 即可。
Sol4
注意到取 \(\min\) 的复杂度是 \(O(n)\) 的。
所以我们采用 ST 表来快速取 \(\min\)。
Code
#include <bits/stdc++.h>
using namespace std;
const int MAXN = 1005;
int n, m;
int lg2[MAXN * 2];
int h[MAXN * 2][MAXN * 2], v[MAXN * 2][MAXN * 2];
int st_h[MAXN * 2][MAXN * 2][12], st_v[MAXN * 2][MAXN * 2][12];
int b[MAXN * 2][MAXN * 2];
vector<int> manacher(vector<int>& s)
{
int len = s.size();
vector<int> rad(len, 1);
for (int i = 0, c = 0, r = 0; i < len; i++)
{
if (i < r) rad[i] = min(r - i, rad[2 * c - i]);
while (i - rad[i] >= 0 && i + rad[i] < len && s[i - rad[i]] == s[i + rad[i]]) rad[i]++;
if (i + rad[i] > r) c = i, r = i + rad[i];
}
return rad;
}
int q_h(int L, int R, int col)
{
int k = lg2[R - L + 1];
return min(st_h[col][L][k], st_h[col][R - (1 << k) + 1][k]);
}
int q_v(int row, int L, int R)
{
int k = lg2[R - L + 1];
return min(st_v[row][L][k], st_v[row][R - (1 << k) + 1][k]);
}
int main()
{
cin >> n >> m;
memset(b, -1, sizeof(b));
for (int i = 1; i <= n; i++)
for (int j = 1; j <= m; j++)
cin >> b[2 * i - 1][2 * j - 1];
for (int i = 2; i < MAXN * 2; i++) lg2[i] = lg2[i / 2] + 1;
for (int i = 1; i <= 2 * n - 1; i++)
{
vector<int> row(2 * m + 1);
for (int j = 0; j <= 2 * m; j++) row[j] = b[i][j];
vector<int> rad = manacher(row);
for (int j = 0; j <= 2 * m; j++) h[i][j] = rad[j];
}
for (int j = 1; j <= 2 * m - 1; j++)
{
vector<int> col(2 * n + 1);
for (int i = 0; i <= 2 * n; i++) col[i] = b[i][j];
vector<int> rad = manacher(col);
for (int i = 0; i <= 2 * n; i++) v[i][j] = rad[i];
}
for (int col = 0; col <= 2 * m; col++)
{
for (int row = 0; row <= 2 * n; row++) st_h[col][row][0] = h[row][col];
for (int k = 1; (1 << k) <= 2 * n + 1; k++)
for (int row = 0; row + (1 << k) - 1 <= 2 * n; row++)
st_h[col][row][k] = min(st_h[col][row][k - 1], st_h[col][row + (1 << (k - 1))][k - 1]);
}
for (int row = 0; row <= 2 * n; row++)
{
for (int col = 0; col <= 2 * m; col++) st_v[row][col][0] = v[row][col];
for (int k = 1; (1 << k) <= 2 * m + 1; k++)
for (int col = 0; col + (1 << k) - 1 <= 2 * m; col++)
st_v[row][col][k] = min(st_v[row][col][k - 1], st_v[row][col + (1 << (k - 1))][k - 1]);
}
long long ans = 0;
for (int i = 0; i <= 2 * n; i++)
{
for (int j = 0; j <= 2 * m; j++)
{
if ((i + j) & 1) continue;
int mx = min(min(i, 2 * n - i), min(j, 2 * m - j));
int jiou = i & 1;
int maxT = jiou ? (mx + 1) / 2 : mx / 2;
int l = 1, r = maxT, best = 0;
while (l <= r)
{
int t = (l + r) >> 1;
int mid = jiou ? (2 * t - 1) : (2 * t);
int rowL = i - mid + 1, rowR = i + mid - 1;
int colL = j - mid + 1, colR = j + mid - 1;
if (q_h(rowL, rowR, j) >= mid && q_v(i, colL, colR) >= mid)
best = t, l = t + 1;
else
r = t - 1;
}
ans += best;
}
}
cout << ans;
return 0;
}

浙公网安备 33010602011771号