把规则和询问都转成字符对串,用 AC 自动机匹配,再在 fail 树上按长度阈值离线计数。
OJ: luogu
题目 ID: P14363
难度:省选/NOI-
标签:字符串AC自动机离线树状数组
日期: 2026-06-22 19:52
题意
给定 n 条替换规则 (s1, s2),保证 |s1| = |s2|。一次替换可以选择原串中的一个子串,如果它等于某条规则的 s1,就把它替换成这条规则的 s2。
每个询问给出两个不同字符串 t1, t2,要求统计有多少种一次替换能把 t1 变成 t2。不同的替换位置或不同的规则编号都算不同方案。
思路
先看一个可以直接验证想法的朴素解:
// 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,必须把所有规则一起匹配。
一次替换不会改变字符串总长度。所以如果 t1 和 t2 长度不同,答案一定是 0。
若长度相同,找到第一个和最后一个不同位置:
L = first position where t1[L] != t2[L]
R = last position where t1[R] != t2[R]合法替换区间必须覆盖 [L,R]。区间外字符不会变化,所以所有不同位置都必须在被替换区间里面。
接着把一位上的变化 (原字符, 目标字符) 看成一个新的字符。比如规则 (s1, s2) 会变成:
(s1[0], s2[0]), (s1[1], s2[1]), ...询问 (t1,t2) 也同样变成字符对串。某条规则能在某个位置完成替换,当且仅当它的字符对串在询问字符对串中出现。
于是可以把所有规则的字符对串插入 AC 自动机。扫描询问到位置 i 时,当前状态的 fail 链上所有终止节点,就是所有以 i 结尾的规则匹配。
还需要满足覆盖 [L,R]。若结束位置是 i,则:
- 必须有
i >= R; - 规则长度至少为
i - L + 1。
所以每个结束位置会产生一个问题:
在当前 AC 状态的 fail 祖先中,统计深度至少为 i-L+1 的终止节点数量。把 fail 指针看成一棵 fail 树。一个终止节点是当前状态的 fail 祖先,等价于当前状态在这个终止节点的子树中。
因此可以离线处理:
- 把所有终止节点按深度从大到小排序;
- 把所有询问拆出的离线问题按需要长度从大到小排序;
- 当终止节点深度足够时,把它的整棵 fail 子树在树状数组中加上该节点的规则数量;
- 查询当前状态的 DFS 序位置,就得到满足长度限制的终止祖先数量。
相同的规则不能去重,因为题目按二元组编号计数。代码用 terminal_weight[u] 记录同一个终止节点上有多少条规则。
代码
#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,需要排序并用树状数组回答。
总时间复杂度可以写作:
O((L1 + L2) log(L1 + L2))空间复杂度为:
O(L1 + L2 + q)总结
本题的关键是把“替换前字符”和“替换后字符”合成一个字符对。这样同时检查左串和右串匹配,就变成了普通多模式串匹配。
覆盖所有不同位置 [L,R] 是另一个关键限制。扫描到每个可能结束位置时,只需要统计长度足够的匹配规则;这个条件可以在 fail 树上按深度阈值离线完成。
一图流解析
这张图把本题的建模、关键转移、实现检查和训练方法压缩到一页,适合读完正文后复盘。
