[CSP-S 2021] 括号序列

GitHub跳转原题关系图返回列表

用区间 DP 分别统计单个外层括号块和由多个块拼接成的合法超级括号序列。

OJ: luogu

题目 ID: P7914

难度:提高+/省选-

标签:动态规划区间 DP字符串

日期: 2026-07-06 08:46

题意

给定长度为 nn 的字符串,每个位置可能是 '('')''*''?'。问有多少种替换所有 '?' 的方式,使得最终字符串是符合定义的“超级括号序列”。答案对 1000000007 取模。

其中连续 '*' 串的长度不能超过 kk,合法序列可以被括号包住,也可以由多个合法部分拼接,中间允许夹一个长度不超过 kk 的星号串。

思路

小数据可以枚举所有 '?' 的替换,再递归判断最终字符串是否合法:

cpp
// brute.cpp:小数据暴力解,枚举所有 ? 的替换,再用递归判断定义是否合法。
#include <bits/stdc++.h>
using namespace std;

const int MOD = 1000000007;

int n, K;
string pattern_s, cur;
vector<int> question_pos;
int memo_valid[20][20], memo_atom[20][20];
long long answer;

bool is_star_string(int l, int r) {
    if (l > r) return false;
    if (r - l + 1 > K) return false;
    for (int i = l; i <= r; i++) {
        if (cur[i] != '*') return false;
    }
    return true;
}

bool is_valid(int l, int r);

bool is_atom(int l, int r) {
    if (l >= r) return false;
    if (memo_atom[l][r] != -1) return memo_atom[l][r];
    bool ok = false;
    if (cur[l] == '(' && cur[r] == ')') {
        if (l + 1 == r) ok = true;
        if (is_star_string(l + 1, r - 1)) ok = true;
        if (l + 1 <= r - 1 && is_valid(l + 1, r - 1)) ok = true;
        for (int p = l + 1; p <= r - 2; p++) {
            if (is_star_string(l + 1, p) && is_valid(p + 1, r - 1)) ok = true;
            if (is_valid(l + 1, p) && is_star_string(p + 1, r - 1)) ok = true;
        }
    }
    memo_atom[l][r] = ok;
    return ok;
}

bool is_valid(int l, int r) {
    if (l > r) return false;
    if (memo_valid[l][r] != -1) return memo_valid[l][r];
    bool ok = is_atom(l, r);
    for (int start = l + 1; start <= r && !ok; start++) {
        if (!is_atom(start, r)) continue;
        if (is_valid(l, start - 1)) ok = true;
        for (int star_len = 1; star_len <= K && start - star_len - 1 >= l; star_len++) {
            int star_l = start - star_len;
            int left_r = star_l - 1;
            if (is_star_string(star_l, start - 1) && is_valid(l, left_r)) {
                ok = true;
            }
        }
    }
    memo_valid[l][r] = ok;
    return ok;
}

void dfs_replace(int idx) {
    if (idx == (int)question_pos.size()) {
        memset(memo_valid, -1, sizeof(memo_valid));
        memset(memo_atom, -1, sizeof(memo_atom));
        if (is_valid(1, n)) {
            answer++;
        }
        return;
    }

    int pos = question_pos[idx];
    cur[pos] = '(';
    dfs_replace(idx + 1);
    cur[pos] = ')';
    dfs_replace(idx + 1);
    cur[pos] = '*';
    dfs_replace(idx + 1);
    cur[pos] = '?';
}

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

    cin >> n >> K >> pattern_s;
    cur = " " + pattern_s;
    for (int i = 1; i <= n; i++) {
        if (cur[i] == '?') {
            question_pos.push_back(i);
        }
    }

    dfs_replace(0);
    cout << (answer % MOD) << '\n';
    return 0;
}

正解使用区间 DP。为了避免拼接规则产生重复计数,区分两类状态:

  • atom[l][r]atom[l][r]:区间 [l,r][l,r] 是一个单个外层括号块
  • f[l][r]f[l][r]:区间 [l,r][l,r] 是一个完整合法超级括号序列,可能由多个 atomatom 拼接而成。

一个 atomatom 必须形如外层一对括号:

text
()
(S)
(A)
(SA)
(AS)

其中 SS 是长度 1k1\dots k 的星号串,AA 是一个合法序列,也就是 ff 状态。

而一个完整序列 f[l][r]f[l][r] 可以是:

  • 单个 atom[l][r]atom[l][r]
  • 前面一段完整序列,加上后面的一个 atomatom
  • 前面一段完整序列,中间隔一个长度不超过 kk 的星号串,再加后面的一个 atomatom

这种“按最后一个 atom 来拆”的方式是唯一的,因此不会把 ABCA B C 这样的串用不同括号方式重复统计。

代码中 sep_sum[l][start]sep\_sum[l][start] 表示:从 ll 开始的一段合法序列,后面接空串或一段合法星号串后,使下一个 atomatomstartstart 开始的方案数。这样转移 f[l][r]f[l][r] 时只需要枚举最后一个 atomatom 的起点 startstart

代码

cpp
// main.cpp:区间 DP,atom 表示单个外层括号块,f 表示若干 atom 拼成的合法序列。
#include <bits/stdc++.h>
using namespace std;

const int MAXN = 505;
const int MOD = 1000000007;

int n, K;
string s;
int bad_star_prefix[MAXN];
int atom_dp[MAXN][MAXN];
int f[MAXN][MAXN];
int sep_sum[MAXN][MAXN]; // sep_sum[l][r]:f[l][..] 后接空串或星串,使下一个 atom 从 r 开始

bool can_left(int pos) {
    return s[pos] == '(' || s[pos] == '?';
}

bool can_right(int pos) {
    return s[pos] == ')' || s[pos] == '?';
}

bool can_star_char(int pos) {
    return s[pos] == '*' || s[pos] == '?';
}

bool is_star_string(int l, int r) {
    if (l > r) {
        return false;
    }
    if (r - l + 1 > K) {
        return false;
    }
    return bad_star_prefix[r] - bad_star_prefix[l - 1] == 0;
}

void add_mod(int &x, long long y) {
    x = (x + y) % MOD;
}

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

    cin >> n >> K;
    cin >> s;
    s = " " + s;

    for (int i = 1; i <= n; i++) {
        bad_star_prefix[i] = bad_star_prefix[i - 1] + (can_star_char(i) ? 0 : 1);
    }

    for (int len = 1; len <= n; len++) {
        for (int l = 1; l + len - 1 <= n; l++) {
            int r = l + len - 1;

            if (l < r) {
                sep_sum[l][r] = f[l][r - 1]; // 两个合法块直接相接
                int max_len = min(K, r - l - 1);
                for (int star_len = 1; star_len <= max_len; star_len++) {
                    int star_l = r - star_len;
                    int left_end = star_l - 1;
                    if (left_end >= l && is_star_string(star_l, r - 1)) {
                        add_mod(sep_sum[l][r], f[l][left_end]);
                    }
                }
            }

            if (len >= 2 && can_left(l) && can_right(r)) {
                if (l + 1 == r) {
                    add_mod(atom_dp[l][r], 1); // ()
                }
                if (is_star_string(l + 1, r - 1)) {
                    add_mod(atom_dp[l][r], 1); // (S)
                }
                if (l + 1 <= r - 1) {
                    add_mod(atom_dp[l][r], f[l + 1][r - 1]); // (A)
                }

                // (SA)
                for (int p = l + 1; p <= r - 2 && p - (l + 1) + 1 <= K; p++) {
                    if (is_star_string(l + 1, p)) {
                        add_mod(atom_dp[l][r], f[p + 1][r - 1]);
                    }
                }

                // (AS)
                for (int p = l + 1; p <= r - 2; p++) {
                    if (is_star_string(p + 1, r - 1)) {
                        add_mod(atom_dp[l][r], f[l + 1][p]);
                    }
                }
            }

            f[l][r] = atom_dp[l][r];
            for (int start = l + 1; start <= r; start++) {
                add_mod(f[l][r], (long long)sep_sum[l][start] * atom_dp[start][r]);
            }
        }
    }

    cout << f[1][n] << '\n';
    return 0;
}

复杂度

区间数是 O(n2)O(n^2)。计算 atomatom 和分隔星号串需要枚举区间内部位置,整体为 O(n3)O(n^3) 级别;n500n \leqslant 500 可以通过。

空间复杂度为 O(n2)O(n^2)

总结

本题最重要的是防止重复计数。把合法串拆成“若干个单独括号块 atom 的序列”,并固定按最后一个 atomatom 来转移,就能把语法规则变成稳定的区间 DP。