QOJ.ac

QOJ

Type: Editorial

Status: Open

Posted by: sjw712

Posted at: 2026-09-12 19:05:50

Last updated: 2026-09-12 19:10:12

Back to Problem

$O(n^{7/2}\sqrt{\log(n+1)})$ by GPT6 Pro

可以做到严格的次四次复杂度:

$$ \boxed{O\!\left(n^{7/2}\sqrt{\log(n+1)}\right)} $$https://qoj.ac/problem/16329/discussion/2665

空间复杂度为 (O(n^3))。这是一个可达到的上界,不代表已经证明它是理论最优复杂度。官方题解给出的是 (O(n^4)) 的逆序贪心 DP;下面是在这个 DP 上推导的“分块 + NTT”优化。(SUA)

代码较长,可以直接取完整文件:C++17 代码

1. 先得到不重复计数的 DP

题目统计的是恢复后的字符串数量,不是划分成 ucup 的方案数,因此不能对同一个 u 的两种用途直接分别计数。

倒序扫描,维护:

$$ a=\#p,\quad b=\#up,\quad h=\#cup,\quad d=\#ucup. $$

其中前三项表示尚未完成的子序列数量。使用如下确定的贪心转移:

读入字符 转移
p (a\gets a+1)
c 要求 (b>0),然后 (b\gets b-1,\ h\gets h+1)
u 若 (a>0),执行 (a\gets a-1,\ b\gets b+1);否则要求 (h>0),执行 (h\gets h-1,\ d\gets d+1)

这也是官方题解采用的判定贪心。(SUA)

关键是 u 优先补 p:假如某个合法划分把当前 u 用于补 cup,那么尚未匹配的 p 必须使用一个更靠左的 u。交换这两个 u 的用途,两个子序列仍然合法。因此优先补 p 不会丢解。

每种字符填法只有一条确定的状态路径,转移系数都是 (1)。

换一组坐标

设当前已经倒序处理了 (i) 个字符,另外记录:

$$ P=\text{已读入的字符 p 总数},\qquad C=\text{已读入的字符 c 总数}. $$

$$ P=a+b+h+d,\qquad C=h+d,\qquad i=a+2b+3h+4d $$

得到

$$ \boxed{ a=2P+C+d-i,\quad b=i-P-2C-d,\quad h=C-d. } $$

于是可以使用滚动数组 dp[d][P][C]。逐字符处理仍然是 (O(n^4)),但这组坐标方便批量转移。

2. 分块:只有边界附近需要逐字符转移

设块长为 (B),当前块实际长度为 (\ell\le B)。

内部状态:(a\ge\ell) 且 (b\ge\ell)

每一步,(a,b) 都最多减少 (1)。因此,从这种状态出发,处理本块的每个字符之前都有 (a,b>0)。

所以本块内:

  • u 一定执行 p → up,不会增加 (d);
  • c 一定可以执行 up → cup
  • 不会触碰贪心的分支边界。

在坐标 ((d,P,C)) 中,这些转移就变成:

$$ u:(d,P,C)\to(d,P,C), $$

$$ p:(d,P,C)\to(d,P+1,C), $$

$$ c:(d,P,C)\to(d,P,C+1). $$

设本块有 (q) 个问号、(f_p) 个固定 p、(f_c) 个固定 c

对于固定的 (d),把块起点的内部状态写成多项式

$$ F_d(x,y)=\sum_{P,C}dp[d][P][C]x^Py^C. $$

整个块的转移就是

$$ \boxed{ G_d(x,y)=F_d(x,y)\,x^{f_p}y^{f_c}(1+x+y)^q. } $$

其中 u 对应 (1),p 对应 (x),c 对应 (y)。卷积核为

$$ [x^r y^s](1+x+y)^q = \frac{q!}{r!\,s!\,(q-r-s)!}. $$

用 NTT 完成二维卷积即可。代码通过把 ((P,C)) 编码成指数 (P W+C),将它转为一次一维卷积;取 (W>n+q),保证第二维不会向第一维进位。

边界状态:(a<\ell) 或 (b<\ell)

这一部分仍按原贪心逐字符 DP。

关键是:从这些状态出发,在本块内始终满足

$$ a<2B\quad\text{或}\quad b<2B. $$

在某一时刻固定 (d,C),上述任意一个条件都只允许 (O(B)) 个 (P)。所以边界部分每层只有

$$ O(n^2B) $$

个可能的状态。

实现时必须只遍历、清空活跃状态,不能每个字符都清空整个三维数组,否则又会退回 (O(n^4))。

3. 复杂度

共有 (O(n/B)) 个块。

内部状态每块对 (O(n)) 个 (d) 分别做一个大小 (O(n^2)) 的卷积,总计

$$ O\!\left(\frac{n^4\log(n+1)}B\right). $$

边界状态逐字符处理,总计

$$ O(n^3B). $$

因此

$$ T(n)=O\!\left( \frac{n^4\log(n+1)}B+n^3B \right). $$

$$ B=\Theta\!\left(\sqrt{n\log(n+1)}\right), $$

得到

$$ \boxed{ T(n)=O\!\left(n^{7/2}\sqrt{\log(n+1)}\right)=o(n^4). } $$

两个部分按块起点状态划分,互不重叠,也没有漏掉状态;内部使用卷积精确统计,边界保留原字符顺序逐步处理。

4. C++17 代码

本地通过了题面全部 6 组样例,并做了小规模合法串枚举对照和覆盖到 (n=50) 的随机对拍;未在在线评测机提交。

#include <bits/stdc++.h>
using namespace std;

constexpr int MOD = 998244353;
constexpr int G = 3;

inline void addmod(int &x, int y) {
    x += y;
    if (x >= MOD) x -= MOD;
}

int modpow(int a, int e) {
    int r = 1;
    for (; e; e >>= 1, a = int(1LL * a * a % MOD))
        if (e & 1) r = int(1LL * r * a % MOD);
    return r;
}

class NTT {
    int n, invN;
    vector<int> rev, roots, invRoots;

public:
    explicit NTT(int size)
        : n(size), invN(modpow(size, MOD - 2)),
          rev(size), roots(size), invRoots(size) {
        for (int i = 1; i < n; ++i)
            rev[i] = (rev[i >> 1] >> 1)
                   | ((i & 1) ? n / 2 : 0);

        int w = modpow(G, (MOD - 1) / n);
        int iw = modpow(w, MOD - 2);

        roots[0] = invRoots[0] = 1;
        for (int i = 1; i < n; ++i) {
            roots[i] = int(1LL * roots[i - 1] * w % MOD);
            invRoots[i] = int(1LL * invRoots[i - 1] * iw % MOD);
        }
    }

    void transform(vector<int> &a, bool inverse) const {
        for (int i = 0; i < n; ++i)
            if (i < rev[i]) swap(a[i], a[rev[i]]);

        const auto &ws = inverse ? invRoots : roots;

        for (int len = 2; len <= n; len <<= 1) {
            int half = len >> 1;
            int step = n / len;

            for (int base = 0; base < n; base += len) {
                for (int j = 0; j < half; ++j) {
                    int x = a[base + j];
                    int y = int(
                        1LL * a[base + j + half]
                        * ws[j * step] % MOD
                    );

                    int sum = x + y;
                    if (sum >= MOD) sum -= MOD;

                    int dif = x - y;
                    if (dif < 0) dif += MOD;

                    a[base + j] = sum;
                    a[base + j + half] = dif;
                }
            }
        }

        if (inverse)
            for (int i = 0; i < n; ++i)
                a[i] = int(1LL * a[i] * invN % MOD);
    }
};

int solve(int n, string s) {
    reverse(s.begin(), s.end());

    const int m = n + 1;
    const int m2 = m * m;
    const int states = m2 * m;
    const int N = 4 * n;

    int B = max(1, (int)ceil(sqrt(n * log2(n + 1.0))));

    vector<int> fact(N + 1, 1), invFact(N + 1, 1);
    for (int i = 1; i <= N; ++i)
        fact[i] = int(1LL * fact[i - 1] * i % MOD);

    invFact[N] = modpow(fact[N], MOD - 2);
    for (int i = N; i >= 1; --i)
        invFact[i - 1] = int(1LL * invFact[i] * i % MOD);

    // dp[d][P][C]:已完成 d 组,已经读入 P 个 p、C 个 c。
    // 展平下标为 d*m2 + P*m + C。
    vector<int> dp(states), result(states);

    // 边界部分使用独立滚动数组和活跃下标。
    vector<int> cur(states), nxt(states), mark(states);
    vector<int> active, nextActive;

    int epoch = 0;
    dp[0] = 1;

    for (int t = 0; t < N; t += B) {
        int len = min(B, N - t);

        int q = 0, fixedP = 0, fixedC = 0;
        for (int j = t; j < t + len; ++j) {
            q += (s[j] == '?');
            fixedP += (s[j] == 'p');
            fixedC += (s[j] == 'c');
        }

        // 整个三维数组只在每个块开始时清空。
        fill(result.begin(), result.end(), 0);
        active.clear();

        // 把 (P,C) 编码为 P*W+C。
        // 卷积中 C 的最大值为 n+q,故 W=n+q+1 不会进位。
        int W = m + q;
        int fftSize = 1;
        while (fftSize < W * W) fftSize <<= 1;

        NTT ntt(fftSize);
        vector<int> kernel(fftSize), poly(fftSize);

        if (q > 0) {
            for (int x = 0; x <= q; ++x) {
                for (int y = 0; x + y <= q; ++y) {
                    int v = int(
                        1LL * fact[q] * invFact[x] % MOD
                    );
                    v = int(1LL * v * invFact[y] % MOD);
                    v = int(
                        1LL * v * invFact[q - x - y] % MOD
                    );
                    kernel[x * W + y] = v;
                }
            }
            ntt.transform(kernel, false);
        }

        for (int d = 0; d <= n; ++d) {
            fill(poly.begin(), poly.end(), 0);
            bool haveInterior = false;

            for (int P = d; P <= n; ++P) {
                for (int C = d; C <= P; ++C) {
                    int id = d * m2 + P * m + C;
                    int val = dp[id];
                    if (val == 0) continue;

                    int a = 2 * P + C + d - t;
                    int b = t - P - 2 * C - d;

                    if (a < len || b < len) {
                        // 边界部分,稍后逐字符转移。
                        cur[id] = val;
                        active.push_back(id);
                    } else if (q == 0) {
                        // 无问号:内部状态直接整体平移。
                        int np = P + fixedP;
                        int nc = C + fixedC;
                        if (np <= n && nc <= n)
                            addmod(
                                result[d * m2 + np * m + nc],
                                val
                            );
                    } else {
                        // 内部部分,放入卷积。
                        haveInterior = true;
                        poly[P * W + C] = val;
                    }
                }
            }

            if (!haveInterior) continue;

            ntt.transform(poly, false);
            for (int k = 0; k < fftSize; ++k)
                poly[k] = int(
                    1LL * poly[k] * kernel[k] % MOD
                );
            ntt.transform(poly, true);

            int end = t + len;
            for (int P = max(d, fixedP); P <= n; ++P) {
                for (int C = max(d, fixedC); C <= P; ++C) {
                    int a = 2 * P + C + d - end;
                    int b = end - P - 2 * C - d;
                    if (a < 0 || b < 0) continue;

                    int v = poly[
                        (P - fixedP) * W + C - fixedC
                    ];
                    addmod(result[d * m2 + P * m + C], v);
                }
            }
        }

        // 只有块起点附近的边界状态逐字符转移。
        // 只清理访问过的格子,不逐字符清空整个数组。
        for (int i = t; i < t + len; ++i) {
            ++epoch;
            nextActive.clear();

            bool canP = s[i] == '?' || s[i] == 'p';
            bool canC = s[i] == '?' || s[i] == 'c';
            bool canU = s[i] == '?' || s[i] == 'u';

            auto put = [&](int id, int val) {
                if (mark[id] != epoch) {
                    mark[id] = epoch;
                    nextActive.push_back(id);
                }
                addmod(nxt[id], val);
            };

            for (int id : active) {
                int val = cur[id];
                cur[id] = 0;
                if (val == 0) continue;

                int d = id / m2;
                int P = id / m % m;
                int C = id % m;

                int a = 2 * P + C + d - i;
                int b = i - P - 2 * C - d;

                if (canP && P < n)
                    put(id + m, val);       // P 增加 1

                if (canC && b > 0)
                    put(id + 1, val);       // C 增加 1

                if (canU) {
                    if (a > 0)
                        put(id, val);       // p -> up
                    else if (C > d)
                        put(id + m2, val);  // cup -> ucup
                }
            }

            cur.swap(nxt);
            active.swap(nextActive);
        }

        // 合并边界部分与内部部分。
        for (int id : active) {
            addmod(result[id], cur[id]);
            cur[id] = 0;
        }

        dp.swap(result);
    }

    return dp[n * m2 + n * m + n];
}

int main() {
    ios::sync_with_stdio(false);
    cin.tie(nullptr);

    int n;
    string s;
    if (!(cin >> n >> s)) return 0;

    cout << solve(n, s) << '\n';
    return 0;
}

Comments

No comments yet.