把所有 DNA 串拼接后按长度 k 的前缀分组,统计每个相同碱基串在各物种中的出现次数,再做组合计数。
OJ: luogu
题目 ID: P8643
难度:提高+/省选-
标签:字符串后缀数组计数组合计数
日期: 2026-06-21 14:08
题意
每个物种有一个 DNA 串,只包含 A/G/C/T。
我们要统计多少个 2m 元组:
(i_1,p_1,i_2,p_2,...,i_m,p_m)
满足:
1 <= i_1 < i_2 < ... < i_m <= n- 从每个物种
i_t的位置p_t开始取出一个长度为k的连续子串 - 这
m个长度为k的子串完全相同
也就是说,要统计“选出 m 个不同物种,并在每个物种中选一个出现位置,使得取出的长度 k 碱基串相同”的方案数。
思路
先看一个可以直接验证想法的朴素解:
#include <bits/stdc++.h>
using namespace std;
int n, m, k;
string s[10];
int choose_id[10];
int choose_pos[10];
long long answer;
void check_current() {
string base = s[choose_id[1]].substr(choose_pos[1], k);
for (int i = 2; i <= m; i++) {
if (s[choose_id[i]].substr(choose_pos[i], k) != base) {
return;
}
}
answer++;
}
void dfs_pos(int dep) {
if (dep > m) {
check_current();
return;
}
int id = choose_id[dep];
int limit = (int)s[id].size() - k;
for (int pos = 0; pos <= limit; pos++) {
choose_pos[dep] = pos;
dfs_pos(dep + 1);
}
}
void dfs_species(int dep, int last) {
if (dep > m) {
dfs_pos(1);
return;
}
for (int i = last + 1; i <= n; i++) {
if ((int)s[i].size() < k) {
continue;
}
choose_id[dep] = i;
dfs_species(dep + 1, i);
}
}
int main() {
ios::sync_with_stdio(false);
cin.tie(nullptr);
cin >> n >> m >> k;
for (int i = 1; i <= n; i++) {
cin >> s[i];
}
// brute.cpp:先枚举选哪些物种,再枚举每个物种里长度为 k 的子串起点,最后检查是否完全相同。
// 复杂度非常高,只适合小数据对拍。
answer = 0;
dfs_species(1, 0);
cout << answer << '\n';
return 0;
}暴力做法就是:
- 枚举选哪
m个物种 - 枚举每个物种里长度为
k的子串起点 - 检查这
m个子串是否完全相同
这个思路很直观,但显然只能做极小数据。
关键是换一个计数顺序。
不要按“选哪些物种和位置”来枚举,而是按“长度为 k 的碱基串内容”来分组统计。
假设某个长度为 k 的串 X:
- 在第 1 个物种里出现了
c_1次 - 在第 2 个物种里出现了
c_2次 - …
- 在第
n个物种里出现了c_n次
那么它对答案的贡献就是:
从 n 个物种里选出 m 个不同物种,然后在每个被选中的物种里任选一个 X 的出现位置。
也就是:
sum(c_{i_1} * c_{i_2} * ... * c_{i_m})
其中 i_1 < i_2 < ... < i_m。
这个值可以用一个很小的组合 DP 求出来,因为题目里 n,m <= 5。
所以真正的问题变成:
如何把所有相同的长度为
k的子串分到同一组里?
这里用后缀数组。
把所有 DNA 串拼接起来,并在每个串后面加一个不同的分隔符,这样就不会跨物种匹配。
然后对整个拼接串建后缀数组和 height 数组。
如果若干个后缀在后缀数组里连续,并且相邻后缀的 LCP >= k,那么它们前 k 个字符就是同一个长度为 k 的子串。
于是每个“相同长度 k 子串”的出现位置,都会对应后缀数组中的一个连续段。
对每个连续段:
- 统计它在每个物种里出现了多少次
- 用
dp[j]表示已经处理了一些物种后,选出j个物种的贡献和 - 转移:
dp[j] += dp[j-1] * cnt[i]
最后 dp[m] 就是这个长度为 k 的子串对答案的贡献。
把所有连续段的贡献加起来即可。
代码
#include <bits/stdc++.h>
using namespace std;
const int MAXN = 100005 + 10;
const int MAXL = 100000 + 5;
const int MAXS = 100000 + 5;
int n, m, k;
string dna[MAXN];
int total_len;
int a[MAXS * 2]; // 拼接后的整数串
int owner[MAXS * 2]; // owner[i] 表示位置 i 属于哪个物种
int valid_pos[MAXS * 2]; // valid_pos[i]=1 表示从这里开始至少还能取出一个长度 k 的子串
int sa[MAXS * 2], rk[MAXS * 2], old_rk[MAXS * 2], tmp_sa[MAXS * 2], cnt_sort[MAXS * 2];
int height_arr[MAXS * 2];
long long dp[10]; // n,m <= 5,这里只要很小的组合 dp
int occur[10]; // 当前长度 k 子串在每个物种里出现多少次
int trans(char ch) {
if (ch == 'A') return 1;
if (ch == 'C') return 2;
if (ch == 'G') return 3;
return 4;
}
void build_sequence() {
total_len = 0;
for (int i = 1; i <= n; i++) {
int len = (int)dna[i].size();
for (int j = 0; j < len; j++) {
total_len++;
a[total_len] = trans(dna[i][j]);
owner[total_len] = i;
valid_pos[total_len] = (j + k <= len);
}
total_len++;
a[total_len] = 4 + i; // 每个物种后面加一个不同分隔符,防止跨串匹配
owner[total_len] = 0;
valid_pos[total_len] = 0;
}
}
void build_sa() {
int maxv = 4 + n;
for (int i = 1; i <= total_len; i++) {
rk[i] = a[i];
sa[i] = i;
}
for (int w = 1;; w <<= 1) {
for (int i = 1; i <= total_len; i++) {
tmp_sa[i] = i;
}
auto cmp = [&](int x, int y) {
if (rk[x] != rk[y]) return rk[x] < rk[y];
int rx = (x + w <= total_len ? rk[x + w] : 0);
int ry = (y + w <= total_len ? rk[y + w] : 0);
return rx < ry;
};
sort(tmp_sa + 1, tmp_sa + total_len + 1, cmp);
for (int i = 1; i <= total_len; i++) {
sa[i] = tmp_sa[i];
}
old_rk[sa[1]] = 1;
int classes = 1;
for (int i = 2; i <= total_len; i++) {
int x = sa[i - 1];
int y = sa[i];
int rx1 = rk[x], ry1 = rk[y];
int rx2 = (x + w <= total_len ? rk[x + w] : 0);
int ry2 = (y + w <= total_len ? rk[y + w] : 0);
if (rx1 != ry1 || rx2 != ry2) {
classes++;
}
old_rk[y] = classes;
}
for (int i = 1; i <= total_len; i++) {
rk[i] = old_rk[i];
}
if (classes == total_len) {
break;
}
}
}
void build_height() {
int h = 0;
for (int i = 1; i <= total_len; i++) {
rk[sa[i]] = i;
}
for (int i = 1; i <= total_len; i++) {
if (rk[i] == 1) {
height_arr[1] = 0;
continue;
}
int j = sa[rk[i] - 1];
if (h > 0) {
h--;
}
while (i + h <= total_len && j + h <= total_len && a[i + h] == a[j + h]) {
h++;
}
height_arr[rk[i]] = h;
}
}
long long count_group(int left, int right) {
for (int i = 1; i <= n; i++) {
occur[i] = 0;
}
for (int i = left; i <= right; i++) {
int pos = sa[i];
if (valid_pos[pos]) {
occur[owner[pos]]++;
}
}
for (int i = 0; i <= m; i++) {
dp[i] = 0;
}
dp[0] = 1;
for (int i = 1; i <= n; i++) {
for (int j = m; j >= 1; j--) {
dp[j] += dp[j - 1] * occur[i];
}
}
return dp[m];
}
int main() {
ios::sync_with_stdio(false);
cin.tie(nullptr);
cin >> n >> m >> k;
for (int i = 1; i <= n; i++) {
cin >> dna[i];
}
build_sequence();
build_sa();
build_height();
long long answer = 0;
for (int i = 1; i <= total_len; ) {
if (!valid_pos[sa[i]]) {
i++;
continue;
}
int j = i;
while (j + 1 <= total_len && height_arr[j + 1] >= k) {
j++;
}
answer += count_group(i, j);
i = j + 1;
}
cout << answer << '\n';
return 0;
}复杂度
设所有字符串总长度为 L。
- 建后缀数组:
- 计算
height: - 扫描连续段并做计数:
,其中 n <= 5
所以总时间复杂度是
总结
这题最重要的是把“按元组枚举”改成“按子串内容分组统计”。
一旦改成这个视角,问题就拆成两部分:
- 用后缀数组把所有相同的长度
k子串聚到一起 - 对每一组做一个很小的组合计数
本质上是一道“字符串分组 + 小规模组合 DP”的题。
一图流解析
这张图把本题的建模、关键转移、实现检查和训练方法压缩到一页,适合读完正文后复盘。
