[蓝桥杯 2016 国 AC] 碱基

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

把所有 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 碱基串相同”的方案数。

思路

先看一个可以直接验证想法的朴素解:

cpp
#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;
}

暴力做法就是:

  1. 枚举选哪 m 个物种
  2. 枚举每个物种里长度为 k 的子串起点
  3. 检查这 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 子串”的出现位置,都会对应后缀数组中的一个连续段。

对每个连续段:

  1. 统计它在每个物种里出现了多少次
  2. dp[j] 表示已经处理了一些物种后,选出 j 个物种的贡献和
  3. 转移: dp[j] += dp[j-1] * cnt[i]

最后 dp[m] 就是这个长度为 k 的子串对答案的贡献。

把所有连续段的贡献加起来即可。

代码

cpp
#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

  • 建后缀数组:O(LlogL)O(L log L)
  • 计算 heightO(L)O(L)
  • 扫描连续段并做计数:O(L+n2)O(L + n^2),其中 n <= 5

所以总时间复杂度是 O(LlogL)O(L log L),空间复杂度是 O(L)O(L)

总结

这题最重要的是把“按元组枚举”改成“按子串内容分组统计”。

一旦改成这个视角,问题就拆成两部分:

  • 用后缀数组把所有相同的长度 k 子串聚到一起
  • 对每一组做一个很小的组合计数

本质上是一道“字符串分组 + 小规模组合 DP”的题。

一图流解析

这张图把本题的建模、关键转移、实现检查和训练方法压缩到一页,适合读完正文后复盘。

一图流解析