很久以前,我自己想过这样一道题,记在洛谷里,没有公开。当时题面已经想好了,但我其实也不太会做,更没有确定应该把数据范围开到多少。
今天忽然想起来,现在 AI 已经这么厉害了,不如把这道题拿出来,看看能不能让 AI 解决一下。
所以先说明一下:题目是我以前想的;这篇 blog 的解题思路、推导、代码和文字整理,全部由 AI 生成。 我提出的问题很简单:能不能做到 \(\mathcal{O}(nm)\)?实在不行,有没有比暴力好一些的方法?
这次得到的结果是:有一个可以直接实现的轻重分治加 bitset 解法,也可以借助快速矩阵乘法得到更好的理论上界。至于最想要的线性算法,目前还没有找到。不过,这个目标和 0/1 矩阵乘法之间有一个很直接的联系。
题目与记号
给定一个 \(n\times m\) 的矩阵 \(a\),输出一个同样大小的矩阵 \(b\),其中
也就是把第 \(i\) 行和第 \(j\) 列放在一起,问一共有多少种不同的数。同一个数无论出现多少次,都只算一次。
例如:
第 \(2\) 行的集合是 \(\{2,3,4\}\),第 \(2\) 列的集合是 \(\{2,4\}\),所以 \(b_{2,2}=3\)。
下面记 \(N=nm\),\(d=\min(n,m)\),\(K\) 为整个矩阵中不同数的个数。假定一个输入值可以用常数个机器字表示;讨论算法核心时,先假设数值已经变成了 \(0,\ldots,K-1\) 的编号。
这一步也需要计入复杂度:比较排序可以在 \(\mathcal{O}(N\log N)\) 时间内完成分组;使用合适的哈希表可以做到期望 \(\mathcal{O}(N)\)。如果输入本来就是 \(\mathcal{O}(N)\) 值域内的整数,也可以直接用数组。后面的参考代码使用排序,因此有确定性的 \(\mathcal{O}(N\log N)\) 预处理开销。
先把重复计算找出来
设第 \(i\) 行的数值集合为 \(S_i\),第 \(j\) 列的数值集合为 \(T_j\)。
行、列各有多少种数都很好算。但把两个数量直接相加时,两边共有的数会被算两次,因此每个共有的数都要减去一次:
例子中 \(S_2\cap T_2=\{2,4\}\),所以应该减去 \(2\)。行与列虽然只交于一个格子,它们的数值集合却可以有很多个公共元素。
现在只剩下一个问题:怎样同时算出所有 \(c_{i,j}\)?
改为枚举一个数的贡献
固定一个数值 \(x\),记
并令 \(r_x=|R_x|\)、\(s_x=|C_x|\)。
只要 \(i\in R_x\) 且 \(j\in C_x\),\(x\) 就是这一行与这一列的公共元素。因此,它恰好给 \(R_x\times C_x\) 中的每一个位置贡献 \(1\)。
例如,矩阵里的数值 \(2\) 出现在全部三行以及前两列。它对 \(c\) 的贡献就是:
这也给出了一个直接的算法:
- 按数值分组,求出每个 \(x\) 的去重行表 \(R_x\) 和去重列表 \(C_x\)。
- 把 \(c\) 初始化为零。
- 对每个 \(x\),枚举 \(R_x\times C_x\),把对应的 \(c_{i,j}\) 加一。
- 用容斥还原 \(b\)。
这里必须先对行号、列号去重。某个数在一行出现五次,也只能在 \(R_x\) 里保留一个行号。
这个算法到底多快?
设 \(f_x\) 是 \(x\) 在原矩阵中的出现次数。每个包含 \(x\) 的行或列都至少消耗一次出现,所以
于是算法的时间是
一方面 \(r_x\le n\),另一方面 \(s_x\le m\),所以
因此最坏复杂度是 \(\mathcal{O}(Nd)\),空间是 \(\mathcal{O}(N)\)。
这个界确实可以达到。对于 \(q\times q\) 的循环拉丁方 \(a_{i,j}=(i+j)\bmod q\),每个数都会出现在所有行、所有列里,直接枚举贡献就要做 \(q^3\) 次加一。
但如果全部元素互不相同,每个数只贡献一个位置,核心计算就只有 \(\mathcal{O}(N)\)。所以,比起只看矩阵大小,\(\sum_x r_xs_x\) 更能反映这个算法在具体数据上的工作量。
轻重分治与 bitset
瓶颈来自那些同时涉及很多行、很多列的数。一个自然的办法是:涉及范围小的数继续枚举,涉及范围大的数放在一起批量处理。
选一个整数阈值 \(L\ge 1\):
- 如果 \(\min(r_x,s_x)\le L\),称 \(x\) 为轻值。
- 如果 \(r_x>L\) 且 \(s_x>L\),称 \(x\) 为重值。
这里按去重后的行数、列数划分。一个数即使出现很多次,只要集中在一行,它的贡献仍然容易枚举。
轻值:直接枚举
若 \(r_x\le L\),它的枚举开销至多为 \(Ls_x\);否则一定有 \(s_x\le L\),开销至多为 \(Lr_x\)。将这两类分别求和,得到
所以轻值部分是 \(\mathcal{O}(NL)\)。
重值:一次处理一个机器字
假设有 \(H\) 个重值。由于每个重值都占据超过 \(L\) 个不同的行,而所有行表的总长度不超过 \(N\),有
把这些重值编号,为每一行、每一列建立 bitset。一个位置为 \(1\),表示对应的重值出现在这一行或这一列里。那么按位与之后,留下的恰好就是双方共有的重值。
因此,一对行列的重值贡献为
这里的 popcount 表示统计二进制中 \(1\) 的个数。每个 bit 对应一个独立的数值,所以结果是精确计数。
设一个机器字有 \(w\) 位,并将一个字的按位与、popcount 视为 \(\mathcal{O}(1)\)。在下面的 C++ 实现里,\(w=64\)。
其实不需要同时存下全部 bitset。每次取至多 \(w\) 个重值,为每行、每列各建一个机器字,扫描全部 \(N\) 个位置,将这一批的公共元素数加到 \(c\) 上,然后复用这些掩码处理下一批即可。
这样有 \(\lceil H/w\rceil\) 批,重值部分的时间为
其中,建立掩码时写入的行号、列号总共只有 \(\mathcal{O}(N)\) 个;每批清空掩码的 \(\mathcal{O}(n+m)\) 开销也被扫描矩阵的开销覆盖。除已有的贡献矩阵和分组数据外,只需要 \(\mathcal{O}(n+m)\) 的掩码空间,因此总空间仍然是 \(\mathcal{O}(N)\)。
怎样选择阈值?
两部分合起来是
把枚举轻值的 \(NL\) 和处理重值的 \(N^2/(Lw)\) 配平,有
还要照顾很窄的矩阵:如果这个阈值已经达到 \(d\),直接令 \(L=d\),所有数都是轻值,退回 \(\mathcal{O}(Nd)\) 的枚举算法。
因此可以选
得到核心复杂度
加上代码采用的排序预处理,就是
这也解释了为什么单独使用轻重分治还不够。如果将一个重值对全部 \(N\) 个位置的贡献逐个处理,重值部分会是 \(\mathcal{O}(N^2/L)\);配平后是 \(\mathcal{O}(N\sqrt N)\),对于方阵仍然是三次方。这里的收益来自 bitset 的位并行。
对于 \(q\times q\) 的方阵,核心上界可以写成
在固定的 64 位机器上,它改善了最坏情况下的运算量,但没有改变三次方这个指数。更精细的实际开销是
通常比统一上界更有参考价值。
C++20 参考实现
输入第一行为 \(n,m\),之后是 \(n\) 行、每行 \(m\) 个整数;输出矩阵 \(b\)。数值使用有符号 64 位整数,算法只关心它们是否相等。
代码按“数值、行号、列号”排序。这样同一个数的行号天然有序,可以相邻去重;列号则用时间戳去重。轻值读完这一组就处理,只有重值需要保留行表和列表。
展开完整代码
#include <algorithm>
#include <bit>
#include <cmath>
#include <cstddef>
#include <cstdint>
#include <iostream>
#include <tuple>
#include <utility>
#include <vector>
struct Cell {
std::int64_t value;
int row;
int col;
};
struct Support {
std::vector<int> rows;
std::vector<int> cols;
};
std::vector<int> intersections(
int n, int m, std::vector<Cell> cells,
std::vector<int>& row_count, std::vector<int>& col_count
) {
const std::size_t total = std::size_t(n) * m;
constexpr int word_bits = 64;
const int threshold = std::min(
std::min(n, m),
std::max(1, int(std::ceil(std::sqrt(
static_cast<long double>(total) / word_bits
))))
);
std::sort(cells.begin(), cells.end(), [](const Cell& a, const Cell& b) {
return std::tie(a.value, a.row, a.col)
< std::tie(b.value, b.row, b.col);
});
row_count.assign(n, 0);
col_count.assign(m, 0);
std::vector<int> common(total, 0);
std::vector<std::size_t> last_column(m, total);
std::vector<Support> heavy;
for (std::size_t begin = 0; begin < total; ) {
std::size_t end = begin;
Support current;
while (end < total && cells[end].value == cells[begin].value) {
const int i = cells[end].row;
const int j = cells[end].col;
if (current.rows.empty() || current.rows.back() != i) {
current.rows.push_back(i);
++row_count[i];
}
if (last_column[j] != begin) {
last_column[j] = begin;
current.cols.push_back(j);
++col_count[j];
}
++end;
}
if (std::min(current.rows.size(), current.cols.size())
<= std::size_t(threshold)) {
for (int i : current.rows) {
const std::size_t offset = std::size_t(i) * m;
for (int j : current.cols) {
++common[offset + j];
}
}
} else {
heavy.push_back(std::move(current));
}
begin = end;
}
// The sorted input is no longer needed.
std::vector<Cell>().swap(cells);
std::vector<std::uint64_t> row_mask(n), col_mask(m);
for (std::size_t begin = 0; begin < heavy.size(); begin += word_bits) {
std::fill(row_mask.begin(), row_mask.end(), 0);
std::fill(col_mask.begin(), col_mask.end(), 0);
const std::size_t end = std::min(heavy.size(), begin + word_bits);
for (std::size_t k = begin; k < end; ++k) {
const std::uint64_t bit = std::uint64_t{1} << (k - begin);
for (int i : heavy[k].rows) row_mask[i] |= bit;
for (int j : heavy[k].cols) col_mask[j] |= bit;
}
for (int i = 0; i < n; ++i) {
const std::size_t offset = std::size_t(i) * m;
for (int j = 0; j < m; ++j) {
common[offset + j] += std::popcount(row_mask[i] & col_mask[j]);
}
}
}
return common;
}
int main() {
std::ios::sync_with_stdio(false);
std::cin.tie(nullptr);
int n, m;
if (!(std::cin >> n >> m) || n <= 0 || m <= 0) return 0;
std::vector<Cell> cells;
cells.reserve(std::size_t(n) * m);
for (int i = 0; i < n; ++i) {
for (int j = 0; j < m; ++j) {
std::int64_t value;
std::cin >> value;
cells.push_back({value, i, j});
}
}
std::vector<int> row_count, col_count;
const auto common = intersections(
n, m, std::move(cells), row_count, col_count
);
for (int i = 0; i < n; ++i) {
for (int j = 0; j < m; ++j) {
const std::int64_t answer = std::int64_t(row_count[i])
+ col_count[j] - common[std::size_t(i) * m + j];
std::cout << answer << (j + 1 == m ? '\n' : ' ');
}
}
}
理论上还能更快吗?
bitset 是把一批重值放在同一个机器字里一起算。还可以把所有重值的贡献放进一次矩阵乘法,使用代数算法来减少运算次数。
先看全部数值。构造两个 0/1 矩阵:
这里 \([P]\) 表示命题 \(P\) 成立时为 \(1\),否则为 \(0\)。固定 \(i,j\),乘积中的每一项 \(U_{i,x}V_{x,j}\),恰好在 \(x\) 同时属于这一行与这一列时贡献 \(1\)。因此
\(U\) 的形状是 \(n\times K\),\(V\) 的形状是 \(K\times m\),而且它们各自至多有 \(N\) 个非零元素。这是一个稀疏矩阵乘法问题。
这里用的是普通加法与乘法。Boolean matrix multiplication 只能告诉我们交集是否非空,无法直接得到交集大小。
把轻的部分直接枚举、重的部分交给快速矩阵乘法,是稀疏矩阵乘法中的经典思路,可以参见 Yuster 和 Zwick 的工作 (Yuster & Zwick, 2005)。对这道题,前面已经给出了轻值 \(\mathcal{O}(NL)\)、重值个数 \(H<N/L\) 的证明,所以可以直接得到:
其中 \(\operatorname{MM}(u,v,z)\) 表示计算一个 \(u\times v\) 矩阵与一个 \(v\times z\) 矩阵乘积的代价,包含构造和读取这两个稠密矩阵的开销。实际计算可以用真正的重值个数 \(H\),不足的维度补零。
这个表达式也适用于长方形矩阵;如果 \(L=d\),没有重值,就完全不需要乘法。数值分组的预处理开销仍然要另外加上。
先用方阵乘法得到一个容易推导的界
设 \(n=m=q\),于是 \(N=q^2\)。假设已有一个 \(\mathcal{O}(q^\gamma)\) 的方阵乘法算法,其中 \(2<\gamma<3\)。
重值部分要乘一个 \(q\times H\) 矩阵和一个 \(H\times q\) 矩阵。沿中间维度每 \(q\) 个数切一块,就能变成 \(\lceil H/q\rceil\) 次方阵乘法,再把结果相加。因此总开销至多是
配平第一项与第三项,取
得到
其中 \(q^\gamma\) 被这个界覆盖,因为 \(\gamma<3\)。
文献中的方阵乘法指数已经小于 \(2.372\) (Alman et al., 2024),所以取 \(\gamma=2.372\),便有 \(\mathcal{O}(q^{2.686})\) 的理论算法。
这里特意使用一个可达到、留有余量的指数 \(\gamma\)。通常记作 \(\omega\) 的最优指数是下确界,直接把它写进大 \(\mathcal{O}\) 时,需要保留任意小的额外指数或 \(o(1)\)。
利用矩形乘法,可以得到 \(\mathcal{O}(q^{2.660})\)
重值矩阵本来就是长方形的。使用针对矩形的乘法界,还能比上面的方阵分块更好。下面给出一个可以由文献中的指数界推出的结果,不声称它是这道题的最优上界。
展开矩形乘法的指数计算
记 \(\omega(1,t,1)\) 为 \(q\times q^t\) 与 \(q^t\times q\) 矩阵乘法的指数下确界。
Alman 等人的论文 Table 1 给出了比下面两个数更强的界 (Alman et al., 2024):
矩阵乘法算法可以通过张量积组合:维度的指数按比例相加,运算量的指数也按比例相加。也就是,可以在这两个形状之间使用凸组合上界。
选取 \(0.32\) 和 \(0.68\) 作为比例,中间维度的指数为
相应的乘法指数满足
现在在本题中取 \(L=\lceil q^{0.660}\rceil\),则重值个数不超过 \(q^{1.340}\)。轻值枚举需要 \(\mathcal{O}(q^{2.660})\),重值矩阵乘法也能在这个界内完成。构造矩阵至多需要 \(\mathcal{O}(q^{2.340})\),同样被覆盖。
因此,整个算法可以做到 \(\mathcal{O}(q^{2.660})\)。上面保留的严格余量允许吸收矩阵乘法指数定义中的任意小额外指数。
以上是代数运算次数的界。要精确计算,可以选取素数 \(d<p\le 2d\),在模 \(p\) 的有限域中完成乘法:每个真实交集大小都不超过 \(d\),所以最终的模 \(p\) 结果就是原来的整数计数。这里仍假设这些 \(\mathcal{O}(\log N)\) 位整数的域运算按常数时间计费。
这些指数界主要回答“理论上能做到多快”。它们使用的快速矩阵乘法构造不适合直接照搬成普通竞赛代码;前面的 C++ 实现采用的仍然是轻重分治与 64 位掩码。
如果真的找到 \(\mathcal{O}(nm)\),会发生什么?
知道“这道题能写成矩阵乘法”,还不足以说明它难。也许这里的矩阵有特殊结构,刚好可以更快处理。
要说明线性算法的意义,需要反过来:把任意两个 0/1 矩阵的乘法,变成这道题。
给定两个 \(k\times k\) 的 0/1 矩阵 \(P,Q\),我们想算普通整数乘积
构造一个 \(2k\times2k\) 的矩阵 \(a\),初始全部填 \(0\)。再用数值 \(1,\ldots,k\) 分别表示乘法的中间下标 \(t\):
- 若 \(P_{i,t}=1\),在右上块的 \(a_{i,k+t}\) 填入 \(t\)。
- 若 \(Q_{t,j}=1\),在左下块的 \(a_{k+t,j}\) 填入 \(t\)。
- 其他位置保持为 \(0\)。
其结构为
这里只用到了 \(0,\ldots,k\) 这 \(k+1\) 种整数。关注答案的左上块:第 \(i\) 行的集合是
第 \(j\) 列的集合是
因为左上块全部是零,\(0\) 一定在两边。除去它,双方共同拥有的数值 \(t\),与乘积里贡献 \(1\) 的中间下标一一对应。因此
记 \(u_i=\sum_tP_{i,t}\)、\(v_j=\sum_tQ_{t,j}\),就有
从而
构造输入、统计 \(u,v\) 和还原全部乘积都只要 \(\mathcal{O}(k^2)\)。如果本题存在通用的 \(\mathcal{O}(nm)\) 算法,那么把它运行在这个 \(2k\times2k\) 的矩阵上,就得到一个 \(\mathcal{O}(k^2)\) 的 0/1 矩阵整数乘法算法。进一步判断每个结果是否大于零,也就得到了 Boolean matrix multiplication。
例如,取
构造出的输入是
它的 \(b\) 的左上块为 \(\begin{pmatrix}3&3\\3&2\end{pmatrix}\)。代入上式,可以还原
这个归约说明,线性算法会带来一个很强的矩阵乘法结果。它没有证明线性算法不存在,也没有给本题建立超过 \(\Omega(nm)\) 的无条件下界。 目前这次探索只能说没有找到通用线性解法,不能把“没有找到”写成“不可能”。
那么数据范围可以开多大?
这道题没有既定的数据范围,最好先看矩阵形状和数值分布,再决定规模。下面的复杂度均不含数值分组的预处理:
| 方法或条件 | 时间复杂度 | 说明 |
|---|---|---|
| 按数值枚举贡献 | \(\mathcal{O}(N+\sum_xr_xs_x)\) | 最坏 \(\mathcal{O}(Nd)\),空间 \(\mathcal{O}(N)\) |
| 轻重分治加 bitset | \(\mathcal{O}(N[1+\min(d,\sqrt{N/w})])\) | 参考代码采用此方法,空间 \(\mathcal{O}(N)\) |
| 方阵:轻重分治加快速矩形乘法 | \(\mathcal{O}(q^{2.660})\) | \(n=m=q\),理论上的代数运算界 |
| \(d=\mathcal{O}(1)\) | \(\mathcal{O}(N)\) | 直接枚举贡献即可 |
| 每个数的 \(\min(r_x,s_x)=\mathcal{O}(1)\) | \(\mathcal{O}(N)\) | 与轻值部分相同的求和证明 |
| 总共只有常数种数 | \(\mathcal{O}(N)\) | 对每种数扫描一次所有位置 |
对于参考实现,估算工作量时可以直接计算轻值的 \(\sum r_xs_x\),再加上 \(N\lceil H/64\rceil\) 次按位与和 popcount。全部不同、全部相同、循环拉丁方,以及大量数值在多行多列中零散出现,都会产生很不一样的开销。只用随机小值域数据测试,很容易漏掉较难的分布。
这次也实际跑了上面的代码。在 Intel Xeon Platinum 8462Y+ 上,使用 GCC 9.3、-std=c++2a -O3 -march=native 编译,每组运行三次,得到下面的中位数。c++2a 是这个版本 GCC 使用的 C++20 模式名称。
| 矩阵大小 | 数值分布 | 排序、分组与计数用时 |
|---|---|---|
| \(1024\times1024\) | 全部相同 | 0.026 s |
| \(1024\times1024\) | 全部不同 | 0.037 s |
| \(1024\times1024\) | 循环拉丁方 | 0.089 s |
| \(1024\times1024\) | 在 8192 种数中均匀随机取值 | 0.201 s |
| \(2048\times2048\) | 在 16384 种数中均匀随机取值 | 1.576 s |
| \(16\times65536\) | 在 128 种数中均匀随机取值 | 0.138 s |
| \(65536\times16\) | 在 128 种数中均匀随机取值 | 0.111 s |
这里不包含生成数据、读入和打印答案的时间,随机数据使用固定种子 20260913。这些是本次环境中的样本测量,不是最坏用时保证;不过至少说明,百万乃至数百万个格子的输入,是这个实现可以实际尝试的量级。
正确性方面,从文章中提取出的代码通过了 33,685 个矩阵的独立集合暴力校验。其中包括 \(1\le n,m\le3\) 以及 \(2\times4\)、\(4\times2\) 矩阵的全部相等关系、随机矩阵、重值数量跨越 64 位分批边界的构造,以及 1,260 组矩阵乘法归约。校验同时开启了 AddressSanitizer 和 UndefinedBehaviorSanitizer。较大数据另外抽样核对了行列并集。
这些结果可以作为设定数据范围的起点。真要出成一道题,还需要针对预期时限、内存和输出量继续测试。尤其要分别考虑方阵和很窄的长方形矩阵,不能只限制 \(nm\) 就认为它们一样难。
这次还留下两个值得继续想的问题:有没有比这里的 bitset 实现更好的实用算法,以及能否利用输入矩阵的结构改进理论上界。最开始希望的 \(\mathcal{O}(nm)\),也仍然是一个没有在这里解决的问题。