[CSP-S 2025] 谐音替换

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

把规则和询问都转成字符对串,用 AC 自动机匹配,再在 fail 树上按长度阈值离线计数。

OJ: luogu

题目 ID: P14363

难度:省选/NOI-

标签:字符串AC自动机离线树状数组

日期: 2026-06-22 19:52

题意

给定 n 条替换规则 (s1, s2),保证 |s1| = |s2|。一次替换可以选择原串中的一个子串,如果它等于某条规则的 s1,就把它替换成这条规则的 s2

每个询问给出两个不同字符串 t1, t2,要求统计有多少种一次替换能把 t1 变成 t2。不同的替换位置或不同的规则编号都算不同方案。

思路

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

cpp
// brute.cpp:小数据暴力解,用来帮助理解题意并辅助对拍。
#include <bits/stdc++.h>
using namespace std;

const int MAXN = 105;

int n, q;
string left_part[MAXN], right_part[MAXN];

bool can_replace_at(const string &a, const string &b, int id, int start) {
    int len = (int)left_part[id].size();
    if (start + len > (int)a.size()) {
        return false;
    }

    for (int i = 0; i < len; i++) {
        if (a[start + i] != left_part[id][i]) {
            return false;
        }
    }

    for (int i = 0; i < (int)a.size(); i++) {
        char after_char = a[i];
        if (start <= i && i < start + len) {
            after_char = right_part[id][i - start];
        }
        if (after_char != b[i]) {
            return false;
        }
    }

    return true;
}

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

    cin >> n >> q;
    for (int i = 1; i <= n; i++) {
        cin >> left_part[i] >> right_part[i];
    }

    while (q--) {
        string a, b;
        cin >> a >> b;

        long long ans = 0;
        if (a.size() == b.size()) {
            for (int id = 1; id <= n; id++) {
                for (int start = 0; start + (int)left_part[id].size() <= (int)a.size(); start++) {
                    if (can_replace_at(a, b, id, start)) {
                        ans++;
                    }
                }
            }
        }
        cout << ans << '\n';
    }

    return 0;
}

暴力会枚举每条规则和每个起点,检查替换后是否等于目标串。它很直观,但规则总长度和询问总长度都可达 5 * 10^6,必须把所有规则一起匹配。

一次替换不会改变字符串总长度。所以如果 t1t2 长度不同,答案一定是 0

若长度相同,找到第一个和最后一个不同位置:

text
L = first position where t1[L] != t2[L]
R = last  position where t1[R] != t2[R]

合法替换区间必须覆盖 [L,R]。区间外字符不会变化,所以所有不同位置都必须在被替换区间里面。

接着把一位上的变化 (原字符, 目标字符) 看成一个新的字符。比如规则 (s1, s2) 会变成:

text
(s1[0], s2[0]), (s1[1], s2[1]), ...

询问 (t1,t2) 也同样变成字符对串。某条规则能在某个位置完成替换,当且仅当它的字符对串在询问字符对串中出现。

于是可以把所有规则的字符对串插入 AC 自动机。扫描询问到位置 i 时,当前状态的 fail 链上所有终止节点,就是所有以 i 结尾的规则匹配。

还需要满足覆盖 [L,R]。若结束位置是 i,则:

  • 必须有 i >= R
  • 规则长度至少为 i - L + 1

所以每个结束位置会产生一个问题:

text
在当前 AC 状态的 fail 祖先中,统计深度至少为 i-L+1 的终止节点数量。

把 fail 指针看成一棵 fail 树。一个终止节点是当前状态的 fail 祖先,等价于当前状态在这个终止节点的子树中。

因此可以离线处理:

  1. 把所有终止节点按深度从大到小排序;
  2. 把所有询问拆出的离线问题按需要长度从大到小排序;
  3. 当终止节点深度足够时,把它的整棵 fail 子树在树状数组中加上该节点的规则数量;
  4. 查询当前状态的 DFS 序位置,就得到满足长度限制的终止祖先数量。

相同的规则不能去重,因为题目按二元组编号计数。代码用 terminal_weight[u] 记录同一个终止节点上有多少条规则。

代码

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

struct TrieEdge {
    int ch;
    int to;
    int next;
};

struct OutEdge {
    int ch;
    int to;
};

struct TreeEdge {
    int to;
    int next;
};

struct Ask {
    int need;
    int state;
    int id;
};

int n, q;

vector<int> trie_head;
vector<int> fail_node;
vector<int> node_depth;
vector<int> out_degree;
vector<long long> terminal_weight;
vector<TrieEdge> trie_edges;

vector<int> edge_start;
vector<OutEdge> sorted_edges;

vector<int> tree_head;
vector<TreeEdge> tree_edges;
vector<int> tin, tout;
int dfs_timer;

vector<int> terminal_nodes;
vector<Ask> asks;
vector<long long> answer;
vector<long long> bit;

int code_pair(char a, char b) {
    return (a - 'a') * 26 + (b - 'a');
}

int new_node() {
    trie_head.push_back(-1);
    fail_node.push_back(0);
    node_depth.push_back(0);
    out_degree.push_back(0);
    terminal_weight.push_back(0);
    return (int)trie_head.size() - 1;
}

int find_child_raw(int u, int ch) {
    for (int e = trie_head[u]; e != -1; e = trie_edges[e].next) {
        if (trie_edges[e].ch == ch) {
            return trie_edges[e].to;
        }
    }
    return -1;
}

int add_child(int u, int ch) {
    int v = new_node();
    node_depth[v] = node_depth[u] + 1;

    TrieEdge e;
    e.ch = ch;
    e.to = v;
    e.next = trie_head[u];
    trie_head[u] = (int)trie_edges.size();
    trie_edges.push_back(e);
    out_degree[u]++;

    return v;
}

int get_or_add_child(int u, int ch) {
    int v = find_child_raw(u, ch);
    if (v != -1) {
        return v;
    }
    return add_child(u, ch);
}

void insert_pair_string(const string &a, const string &b) {
    int u = 0;
    for (int i = 0; i < (int)a.size(); i++) {
        int ch = code_pair(a[i], b[i]);
        u = get_or_add_child(u, ch);
    }
    terminal_weight[u]++;
}

bool cmp_out_edge(const OutEdge &a, const OutEdge &b) {
    return a.ch < b.ch;
}

void build_sorted_edges() {
    int nodes = (int)trie_head.size();
    edge_start.assign(nodes + 1, 0);
    for (int i = 0; i < nodes; i++) {
        edge_start[i + 1] = edge_start[i] + out_degree[i];
    }

    sorted_edges.resize(trie_edges.size());
    vector<int> cur = edge_start;
    for (int u = 0; u < nodes; u++) {
        for (int e = trie_head[u]; e != -1; e = trie_edges[e].next) {
            int pos = cur[u]++;
            sorted_edges[pos].ch = trie_edges[e].ch;
            sorted_edges[pos].to = trie_edges[e].to;
        }
    }

    for (int u = 0; u < nodes; u++) {
        int l = edge_start[u];
        int r = edge_start[u + 1];
        if (r - l > 1) {
            sort(sorted_edges.begin() + l, sorted_edges.begin() + r, cmp_out_edge);
        }
    }
}

int find_child(int u, int ch) {
    int l = edge_start[u];
    int r = edge_start[u + 1] - 1;
    while (l <= r) {
        int mid = (l + r) >> 1;
        if (sorted_edges[mid].ch == ch) {
            return sorted_edges[mid].to;
        }
        if (sorted_edges[mid].ch < ch) {
            l = mid + 1;
        } else {
            r = mid - 1;
        }
    }
    return -1;
}

int move_state(int u, int ch) {
    while (u != 0 && find_child(u, ch) == -1) {
        u = fail_node[u];
    }
    int v = find_child(u, ch);
    if (v == -1) {
        return 0;
    }
    return v;
}

void build_ac_automaton() {
    build_sorted_edges();

    vector<int> que;
    que.reserve(trie_head.size());

    for (int e = edge_start[0]; e < edge_start[1]; e++) {
        int v = sorted_edges[e].to;
        fail_node[v] = 0;
        que.push_back(v);
    }

    for (int head = 0; head < (int)que.size(); head++) {
        int u = que[head];
        for (int e = edge_start[u]; e < edge_start[u + 1]; e++) {
            int ch = sorted_edges[e].ch;
            int v = sorted_edges[e].to;

            int f = fail_node[u];
            while (f != 0 && find_child(f, ch) == -1) {
                f = fail_node[f];
            }
            int go = find_child(f, ch);
            if (go == -1) {
                fail_node[v] = 0;
            } else {
                fail_node[v] = go;
            }
            que.push_back(v);
        }
    }
}

void add_tree_edge(int u, int v) {
    TreeEdge e;
    e.to = v;
    e.next = tree_head[u];
    tree_head[u] = (int)tree_edges.size();
    tree_edges.push_back(e);
}

void build_fail_tree() {
    int nodes = (int)trie_head.size();
    tree_head.assign(nodes, -1);
    tree_edges.reserve(nodes - 1);

    for (int i = 1; i < nodes; i++) {
        add_tree_edge(fail_node[i], i);
    }

    tin.assign(nodes, 0);
    tout.assign(nodes, 0);
    dfs_timer = 0;

    vector<int> st;
    vector<int> iter;
    st.push_back(0);
    iter.push_back(tree_head[0]);
    tin[0] = ++dfs_timer;

    while (!st.empty()) {
        int u = st.back();
        int &e = iter.back();
        if (e != -1) {
            int v = tree_edges[e].to;
            e = tree_edges[e].next;
            tin[v] = ++dfs_timer;
            st.push_back(v);
            iter.push_back(tree_head[v]);
        } else {
            tout[u] = dfs_timer;
            st.pop_back();
            iter.pop_back();
        }
    }
}

void bit_add(int pos, long long val) {
    int nbit = (int)bit.size() - 1;
    while (pos <= nbit) {
        bit[pos] += val;
        pos += pos & -pos;
    }
}

void bit_range_add(int l, int r, long long val) {
    bit_add(l, val);
    bit_add(r + 1, -val);
}

long long bit_query(int pos) {
    long long res = 0;
    while (pos > 0) {
        res += bit[pos];
        pos -= pos & -pos;
    }
    return res;
}

bool cmp_terminal_depth(int a, int b) {
    return node_depth[a] > node_depth[b];
}

bool cmp_ask_need(const Ask &a, const Ask &b) {
    return a.need > b.need;
}

void answer_offline_asks() {
    int nodes = (int)trie_head.size();
    bit.assign(nodes + 2, 0);

    sort(terminal_nodes.begin(), terminal_nodes.end(), cmp_terminal_depth);
    sort(asks.begin(), asks.end(), cmp_ask_need);

    int ptr = 0;
    for (int i = 0; i < (int)asks.size(); i++) {
        while (ptr < (int)terminal_nodes.size() && node_depth[terminal_nodes[ptr]] >= asks[i].need) {
            int u = terminal_nodes[ptr];
            bit_range_add(tin[u], tout[u], terminal_weight[u]);
            ptr++;
        }
        answer[asks[i].id] += bit_query(tin[asks[i].state]);
    }
}

void process_query(int id, const string &a, const string &b) {
    if (a.size() != b.size()) {
        answer[id] = 0;
        return;
    }

    int len = (int)a.size();
    int first_diff = -1;
    int last_diff = -1;
    for (int i = 0; i < len; i++) {
        if (a[i] != b[i]) {
            if (first_diff == -1) {
                first_diff = i;
            }
            last_diff = i;
        }
    }

    if (first_diff == -1) {
        answer[id] = 0;
        return;
    }

    int state = 0;
    for (int i = 0; i < len; i++) {
        int ch = code_pair(a[i], b[i]);
        state = move_state(state, ch);
        if (i >= last_diff) {
            Ask ask;
            ask.need = i - first_diff + 1;
            ask.state = state;
            ask.id = id;
            asks.push_back(ask);
        }
    }
}

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

    cin >> n >> q;
    new_node();

    string a, b;
    for (int i = 1; i <= n; i++) {
        cin >> a >> b;
        insert_pair_string(a, b);
    }

    build_ac_automaton();
    build_fail_tree();

    int nodes = (int)trie_head.size();
    for (int i = 1; i < nodes; i++) {
        if (terminal_weight[i] > 0) {
            terminal_nodes.push_back(i);
        }
    }

    answer.assign(q + 1, 0);
    asks.reserve(1000000);
    for (int id = 1; id <= q; id++) {
        cin >> a >> b;
        process_query(id, a, b);
    }

    answer_offline_asks();

    for (int id = 1; id <= q; id++) {
        cout << answer[id] << '\n';
    }

    return 0;
}

复杂度

设规则总长度为 L1,询问总长度为 L2

建 AC 自动机和扫描询问都是按总长度处理。离线问题数量不超过 L2,需要排序并用树状数组回答。

总时间复杂度可以写作:

text
O((L1 + L2) log(L1 + L2))

空间复杂度为:

text
O(L1 + L2 + q)

总结

本题的关键是把“替换前字符”和“替换后字符”合成一个字符对。这样同时检查左串和右串匹配,就变成了普通多模式串匹配。

覆盖所有不同位置 [L,R] 是另一个关键限制。扫描到每个可能结束位置时,只需要统计长度足够的匹配规则;这个条件可以在 fail 树上按深度阈值离线完成。

一图流解析

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

一图流解析