可以用 分块 + 差值排序 + 位集 加速 max-plus 矩阵乘法,得到一个确定性的改进。
单组最坏时间复杂度为
$$ \boxed{ O\!\left( \left(\frac{8^n}{\sqrt n}+n^{3/2}4^n\right)\log R \right) } $$
空间复杂度为 (O(4^n))。这里按通常的 word-RAM 模型计时。这个改进降低的是一个多项式因子,指数底数仍然是 (8),不是 (O(4^n\log R))。
下文用题面中的 (R) 表示回合数;(T) 是测试组数。本题 (n\le 6),所以状态数至多 (64),一个 uint64_t 就能表示全部候选列。
1. 状压和快速幂框架不变
定义
$$ a[S]=\sum_{i\in S}a_i,\qquad c[S]=\sum_{i\in S}c_i. $$
上一回合使用集合 (P),这一回合使用集合 (Q),合法条件为
$$ c[Q]+k\operatorname{popcount}(P\cap Q)\le m. $$
因为只有连续两个回合都使用的角色,本回合需要额外支付 (k)。
只保留 (c[S]\le m) 的集合,设剩下 (N\le 2^n) 个状态。建立转移矩阵
$$ M_{P,Q}= \begin{cases} a[Q],&P\to Q\text{ 合法},\\ -\infty,&\text{否则}. \end{cases} $$
初始只有空集状态为 (0)。最终计算初始向量乘 (M^R),再取最大值。普通实现每次矩阵乘法是 (O(N^3)),这就是原来 (O(8^n\log R)) 的瓶颈。(Universal Cup Judging System)
2. 如何加速矩阵乘法
需要计算
$$ Z_{i,j}=\max_t\bigl(X_{i,t}+Y_{t,j}\bigr). $$
把中间下标 (t) 分成大小约为 (b) 的块。重点是:对一个块,先批量确定每个 ((i,j)) 的最优下标,然后只计算这个下标的贡献。
比较两个候选,可以转成差值比较
对于块内两个下标 (p<q),有
$$ X_{i,p}+Y_{p,j}\ge X_{i,q}+Y_{q,j} \iff X_{i,p}-X_{i,q}\ge Y_{q,j}-Y_{p,j}. $$
记
$$ L_i=X_{i,p}-X_{i,q},\qquad D_j=Y_{q,j}-Y_{p,j}. $$
将所有 (L_i)、所有 (D_j) 分别排序,再双指针扫描,就能求出每一行对应的位集
$$ F_i=\{j\mid D_j\le L_i\}. $$
对这些列,候选 (p) 不差于 (q)。约定相等时较小下标获胜,因此可以更新:
win[i][p] &= F_i;
win[i][q] &= full ^ F_i;
其中 win[i][p] 表示:第 (i) 行中,候选 (p) 仍可能获胜的列。
处理完块内所有候选对以后,对于每个 ((i,j)),恰好只有一个候选的位集中还保留着第 (j) 位:它就是该块内取得最大值的最小下标。
于是,一个块中所有结果的更新总共只有 (N^2) 次,而不是 (N^2b) 次。
3. 复杂度
共有 (O(N/b)) 个块。每块有 (O(b^2)) 个候选对,每对需要排序两个长度为 (N) 的数组。
在本题的单字位集实现中,一次矩阵乘法为
$$ O\left(N^2b\log N+\frac{N^3}{b}\right). $$
为了严格计入位集代价,推广到任意 (N) 时,设机器字长为 (w),则复杂度为
$$ O\left( N^2b\log N+\frac{N^3}{b}+\frac{N^3b}{w} \right). $$
取
$$ b=\Theta(\sqrt{\log N}), $$
在 (w=\Omega(\log N)) 的 word-RAM 模型下,得到
$$ O\left( \frac{N^3}{\sqrt{\log N}} +N^2(\log N)^{3/2} \right). $$
代入 (N\le 2^n),再乘快速幂的 (\log R),就是开头的复杂度。
4. C++17 代码
代码保留了不可达状态处理,没有猜测循环节。相等时统一让较小下标获胜,这一点保证了每个块只有 (N^2) 次有效候选枚举。
#include <bits/stdc++.h>
using namespace std;
using i64 = long long;
using u64 = uint64_t;
constexpr int LIM = 64;
constexpr i64 NEG = -(1LL << 60);
struct Matrix {
int n;
i64 a[LIM][LIM];
explicit Matrix(int n_ = 0) : n(n_) {
for (int i = 0; i < n; ++i)
fill(a[i], a[i] + n, NEG);
}
};
struct Item {
i64 value;
int index;
bool operator<(const Item& other) const {
return value < other.value;
}
};
Matrix multiply(const Matrix& A, const Matrix& B) {
const int N = A.n;
Matrix C(N);
// 本题 N <= 64,一个 uint64_t 表示所有列。
// 分块大小约为 sqrt(log2 N)。
const int b = max(
1, (int)ceil(sqrt(max(1.0, log2((double)N))))
);
const u64 full =
(N == 64 ? ~u64(0) : (u64(1) << N) - 1);
u64 win[LIM][LIM];
Item left[LIM], right[LIM];
for (int st = 0; st < N; st += b) {
const int len = min(b, N - st);
for (int i = 0; i < N; ++i)
fill(win[i], win[i] + len, full);
for (int p = 0; p < len; ++p) {
for (int q = p + 1; q < len; ++q) {
const int u = st + p;
const int v = st + q;
for (int i = 0; i < N; ++i)
left[i] = {
A.a[i][u] - A.a[i][v], i
};
for (int j = 0; j < N; ++j)
right[j] = {
B.a[v][j] - B.a[u][j], j
};
sort(left, left + N);
sort(right, right + N);
int ptr = 0;
u64 mask = 0;
for (int t = 0; t < N; ++t) {
while (ptr < N &&
right[ptr].value <= left[t].value) {
mask |= u64(1) << right[ptr].index;
++ptr;
}
const int i = left[t].index;
// 相等时让较小的下标 u 获胜。
win[i][p] &= mask;
win[i][q] &= full ^ mask;
}
}
}
// 每一行中,所有 win 位集恰好划分全部 N 列。
for (int i = 0; i < N; ++i) {
for (int p = 0; p < len; ++p) {
const int u = st + p;
if (A.a[i][u] == NEG) continue;
u64 mask = win[i][p];
while (mask) {
const int j = __builtin_ctzll(mask);
mask &= mask - 1;
if (B.a[u][j] == NEG) continue;
const i64 value = A.a[i][u] + B.a[u][j];
if (value > C.a[i][j])
C.a[i][j] = value;
}
}
}
}
return C;
}
vector<i64> apply_matrix(const vector<i64>& f,
const Matrix& A) {
const int N = A.n;
vector<i64> g(N, NEG);
for (int i = 0; i < N; ++i) {
if (f[i] == NEG) continue;
for (int j = 0; j < N; ++j) {
if (A.a[i][j] == NEG) continue;
g[j] = max(g[j], f[i] + A.a[i][j]);
}
}
return g;
}
int main() {
ios::sync_with_stdio(false);
cin.tie(nullptr);
int tests;
if (!(cin >> tests)) return 0;
while (tests--) {
int n, m, k;
i64 R;
cin >> n >> m >> k >> R;
vector<int> damage(n), cost(n);
for (int i = 0; i < n; ++i)
cin >> damage[i] >> cost[i];
const int S = 1 << n;
vector<i64> val(S);
vector<int> base(S);
vector<int> masks;
for (int s = 0; s < S; ++s) {
if (s) {
const int bit = __builtin_ctz((unsigned)s);
const int t = s & (s - 1);
val[s] = val[t] + damage[bit];
base[s] = base[t] + cost[bit];
}
if (base[s] <= m)
masks.push_back(s);
}
const int N = (int)masks.size();
Matrix A(N);
for (int i = 0; i < N; ++i) {
for (int j = 0; j < N; ++j) {
const int repeated = __builtin_popcount(
(unsigned)(masks[i] & masks[j])
);
if (base[masks[j]] + k * repeated <= m)
A.a[i][j] = val[masks[j]];
}
}
vector<i64> f(N, NEG);
f[0] = 0; // 第 0 回合使用空集;masks[0] 一定是 0。
while (R > 0) {
if (R & 1)
f = apply_matrix(f, A);
R >>= 1;
if (R)
A = multiply(A, A);
}
cout << *max_element(f.begin(), f.end()) << '\n';
}
return 0;
}
这里用有限的 NEG 做差值比较是安全的:所有真实伤害非负,最大答案不超过 (6\times10^{15}),远小于 (2^{60});只要存在合法候选,它一定胜过不可达候选,真正更新时也会排除不可达项。上述数值界来自题面的伤害和回合数范围。
已通过题面样例、4000 组小回合逐轮 DP 对拍,以及 1100 组大回合普通矩阵快速幂对拍;未在 OJ 提交验证。