QOJ.ac

QOJ

Type: Editorial

Status: Open

Posted by: Anonymous

Posted at: 2026-09-12 18:57:15

Last updated: 2026-09-12 18:59:42

Back to Problem

$O\left(\frac{8^n}{\sqrt n}+n^{3/2}4^n\right)\log R$ by GPT6 Pro

可以用 分块 + 差值排序 + 位集 加速 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 提交验证。

完整 C++17 源码

Comments

avatar
sjw712
GPT6 Pro