题解:P16958 [SCCPC 2026] 括号序列

· · 题解

孩子们来个全题解区最劣解(恼)

下文 n 均指括号串中的括号对数。

枚举一个右括号尝试改成左括号。设 1\sim ia_i 个左括号,c_i 个右括号。枚举后面的长度可以得到答案为

\sum^{n-a_i-1}_m[x^m]F^{a_i-c_i+3}=[x^{n-a_i-1}]\dfrac{F^{a_i-c_i+3}}{1-x}

其中 F 是卡特兰数的 OGF,其满足 F=1+xF^2。设 y=xF,那么有 x=y(1-y)F=(1-y)^{-1},得到答案为:

[x^n]\dfrac{x^{a_i+1}F^{a_i-c_i+3}}{1-x}=[x^n]\dfrac{y^{a_i+1}(1-y)^{c_i-1}}{(1-y)(1-x)}

P(y)=\sum y^{a_i+1}(1-y)^{c_i-1}。由于 a_i,c_i 都是单调的,可以直接分治 NTT 求出来。而分母是 (1-y)(1-x)=(1-y)(1-y-y^2)=1-2y+2y^2-y^3,可以直接做一个三阶的整式递推得到答案关于 y 的多项式。那么只要求出 [x^n]y^k 就赢了。考虑拉反:

[x^n]y^k=\frac{k}{n}[z^{n-k}](1-z)^{-n}=\frac kn\dbinom{2n-k-1}{n-k}

至此本题结束,总时间复杂度 \mathcal{O}(n\log^2n)。因为两只老哥跑的快随便过。

:::info[Code]

#include <bits/stdc++.h>

using namespace std;

typedef long long ll;

const int MAXN = 1e6 + 10;
const int mod = 998244353, rt = 3;

inline 
int qpow(int b, int p) {
    int res = 1;
    for (; p; p >>= 1, b = (ll)b * b % mod) if (p & 1) res = (ll)res * b % mod;
    return res;
}

int urt[1 << 21];

inline 
void init_rt(int l) {
    urt[l >> 1] = 1; int x = qpow(rt, (mod - 1) / l);
    for (int i = (l >> 1) + 1; i < l; i++) urt[i] = (ll)urt[i - 1] * x % mod;
    for (int i = (l >> 1) - 1; i > 0; i--) urt[i] = urt[i << 1];
}

inline 
void dft(vector<int> &f, int l) {
    for (int i = l >> 1; i; i >>= 1) {
        for (int j = 0; j < l; j += i << 1) {
            for (int k = j; k < j + i; k++) {
                int x = f[i + k];
                f[i + k] = (ll)(f[k] - x + mod) * urt[i + k - j] % mod;
                f[k] += x, f[k] < mod || (f[k] -= mod);
            }
        }
    }
}

inline 
void idft(vector<int> &f, int l) {
    for (int i = 1; i < l; i <<= 1) {
        for (int j = 0; j < l; j += i << 1) {
            for (int k = j; k < j + i; k++) {
                int x = (ll)f[i + k] * urt[i + k - j] % mod;
                f[i + k] = f[k] - x, f[i + k] < 0 && (f[i + k] += mod);
                f[k] += x, f[k] < mod || (f[k] -= mod);
            }
        }
    }
    int x = mod - mod / l;
    for(int i = 0; i < l; i++) f[i] = (ll)f[i] * x % mod;
    reverse(f.begin() + 1, f.begin() + l);
}

inline int add(int x, int y) { return x += y, x < mod ? x : x - mod; }
inline int sub(int x, int y) { return x -= y, x < 0 ? x + mod : x; }
inline void cadd(int &x, int y) { x += y, x < mod || (x -= mod); }
inline void csub(int &x, int y) { x -= y, x < 0 && (x += mod); }

int fac[MAXN], ifac[MAXN];

inline 
void init(int n) {
    *fac = 1;
    for (int i = 1; i <= n; i++) fac[i] = (ll)fac[i - 1] * i % mod;
    ifac[n] = qpow(fac[n], mod - 2);
    for (int i = n; i; i--) ifac[i - 1] = (ll)ifac[i] * i % mod;
}

inline 
int C(int n, int m) {
    if (n < 0 || m < 0 || n < m) return 0;
    return (ll)fac[n] * ifac[m] % mod * ifac[n - m] % mod;
}

int T, n, a[MAXN], pos[MAXN]; char s[MAXN];

int p[MAXN];

vector<int> solve(int l, int r, int lim) {
    if (lim < 0) return vector<int>();
    if (r - l == 1) return vector<int>(1, 1);
    if (pos[l] == pos[r - 1]) {
        int d = min(lim, r - l - 1);
        vector<int> f(d + 1);
        for (int i = 0; i <= d; i++) (i & 1 ? csub : cadd)(f[i], C(r - l, i + 1));
        return f;
    }
    int mid = (l + r) >> 1;
    vector<int> L = solve(l, mid, lim);
    int ofs = pos[mid] - pos[l], rlim = lim - ofs;
    vector<int> R = solve(mid, r, rlim);
    if (rlim >= 0 && !R.empty()) {
        int d = min(mid - l, rlim);
        if (d) {
            vector<int> f(d + 1);
            for (int i = 0; i <= d; i++) (i & 1 ? csub : cadd)(f[i], C(mid - l, i));
            int tl = min(rlim + 1, (int)R.size() + d), len = 1;
            for (; len < R.size() + d; len <<= 1);
            R.resize(len), f.resize(len), init_rt(len), dft(R, len), dft(f, len);
            for (int i = 0; i < len; i++) R[i] = (ll)R[i] * f[i] % mod;
            idft(R, len), R.resize(tl);
        }
        int tl = min(lim + 1, max<int>(L.size(), ofs + R.size())); L.resize(tl);
        for (int i = 0; i < R.size() && i + ofs < tl; i++) cadd(L[i + ofs], R[i]);
    }
    if (L.size() > lim + 1) L.resize(lim + 1); return L;
}

int main() {
    freopen("brac.in", "r", stdin);
    freopen("brac.out", "w", stdout);
    init(1e6);
    for (scanf("%*d%d", &T); T--; ) {
        scanf("%d%s", &n, s + 1); int m = n >> 1, ans = 0;
        for (int i = 1; i <= n; i++) a[i] = a[i - 1] + (s[i] == '(');
        for (int i = 0; i <= n; i++) p[i] = 0;
        for (int i = 1, j = 0; i <= n; i++) if (s[i] == ')') pos[j++] = a[i];
        {
            vector<int> f = solve(0, m, m - *pos - 1);
            for (int i = 0; i < f.size(); i++) p[i + *pos + 1] = f[i];
        }
        for (int i = 0; i <= m; i++) {
            if (i > 0) cadd(p[i], add(p[i - 1], p[i - 1]));
            if (i > 1) csub(p[i], add(p[i - 2], p[i - 2]));
            if (i > 2) cadd(p[i], p[i - 3]);
        }
        for (int i = 1; i <= m; i++) {
            int w = (ll)i * C(n - i - 1, m - i) % mod;
            ans = (ans + (ll)p[i] * w) % mod;
        }
        ans = (ll)ans * qpow(m, mod - 2) % mod;
        for (int i = 1, x = 0; i <= n; i++) {
            s[i] == '(' ? ++x : --x;
            if (!x) cadd(ans, 1);
        }
        printf("%d\n", ans); 
    }
}

:::