看到字符串匹配,我们就要往 AC 自动机上去想。我们根据 AC 自动机的性质,我们很容易得知在第 $i$ 位结尾的匹配有多少个,所以我们可以很轻松地得到一个文本串的前缀的匹配数。
因此,我们考虑:假设第 $i$ 个位置的匹配数是 $suc_i$,那么我们能不能把区间 $[l, r]$ 的匹配数转化成 $suc_r - suc_l$ 这样的形式?大体可行,但是有一些情况我们需要处理。
我们询问的区间是 $[l, r]$,有一个匹配成功的字符串 $[x, y]$,如果 $1 \leq x < l \leq y \leq r $ 那么 $[x, y]$ 显然是被多算的。我们要排除这种影响。
定义 $sl_i$ 表示在第 $i$ 位结尾的所有匹配中的起始点(这个值可以用 fail 树随便维护一下),如果 $sl_i < l$,那么存在多算的匹配。如果我们找到 $[l, r]$ 中最大的使得 $sl_i < l$ 的 $i$,记作 $m$,那么在第 $[m + 1, r]$ 位之间结尾的匹配一定不会多算,所以直接加上 $suc_r - suc_m$。$m$ 我们可以离线二分或者在线的线段树二分维护。
接下来处理 $[l, m]$ 之间的匹配。根据刚才的条件,我们一定存在一个模式串,使得 $[l, m]$ 一定是这个模式串的后缀,所以相当于我们要求出来每一个模式串后缀的匹配数。我们把所有模式串的反串扔到一个 AC 自动机上,对于每一个模式串倒着匹配一遍,就可以以 $\sum |s_i|$ 的复杂度求出每一个模式串的后缀匹配数。把这部分的答案加上就好。
代码如下:
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188
| #include <bits/stdc++.h> using namespace std; #ifndef ONLINE_JUDGE #include "my_local.h" #endif #define ll long long inline int read() { int x = 0, f = 1; char ch = getchar(); while (ch < '0' || ch > '9') {if (ch == '-') f = -1; ch = getchar();} while (ch >= '0' && ch <= '9') {x = x * 10 + ch - '0'; ch = getchar();} return x * f; }
struct Trie { int son[26]; int cnt, fail, idx; } tp[1000010], ts[1000010]; int idxp[500010], idxs[500010]; int cntp, cnts; vector<int> Gp[1000010], Gs[1000010]; ll sump[1000010], sums[1000010], suc[5000010]; int dep[1000010], lst[1000010]; int lgt[5000010], pos[5000010]; int ci; string pp[500010]; vector<ll> g[500010]; void insertp(char s[], int &idx) { int len = strlen(s + 1), u = 0; for (int i = 1; i <= len; i++) { int &son = tp[u].son[s[i] - 'a']; if (!son) son = ++cntp; u = son; } tp[u].cnt++; tp[u].idx = ci; idx = u; } void inserts(char s[], int &idx) { int len = strlen(s + 1), u = 0; for (int i = len; i; i--) { int &son = ts[u].son[s[i] - 'a']; if (!son) son = ++cnts; u = son; } ts[u].cnt++; ts[u].idx = ci; idx = u; } void buildp() { queue<int> q; for (int i = 0; i < 26; i++) if (tp[0].son[i]) q.push(tp[0].son[i]), Gp[0].push_back(tp[0].son[i]); while (!q.empty()) { int u = q.front(); q.pop(); for (int i = 0; i < 26; i++) { if (tp[u].son[i]) { tp[tp[u].son[i]].fail = tp[tp[u].fail].son[i]; Gp[tp[tp[u].fail].son[i]].push_back(tp[u].son[i]); q.push(tp[u].son[i]); } else tp[u].son[i] = tp[tp[u].fail].son[i]; } } } void builds() { queue<int> q; for (int i = 0; i < 26; i++) if (ts[0].son[i]) q.push(ts[0].son[i]), Gs[0].push_back(ts[0].son[i]); while (!q.empty()) { int u = q.front(); q.pop(); for (int i = 0; i < 26; i++) { if (ts[u].son[i]) { ts[ts[u].son[i]].fail = ts[ts[u].fail].son[i]; Gs[ts[ts[u].fail].son[i]].push_back(ts[u].son[i]); q.push(ts[u].son[i]); } else ts[u].son[i] = ts[ts[u].fail].son[i]; } } } void dfs(int u) { for (int i = 0; i < 26; i++) { int v = tp[u].son[i]; if (v) { dep[v] = dep[u] + 1; dfs(v); } } } void dfsp(int u) { for (int v : Gp[u]) { sump[v] = sump[u] + tp[v].cnt; lst[v] = lst[u]; if (tp[v].cnt) lst[v] = v; dfsp(v); } } void dfss(int u) { for (int v : Gs[u]) { sums[v] = sums[u] + ts[v].cnt; dfss(v); } } void match(char s[]) { int len = strlen(s + 1), u = 0; for (int i = 1; i <= len; i++) { suc[i] = suc[i - 1]; u = tp[u].son[s[i] - 'a']; pos[i] = u; suc[i] += sump[u]; lgt[i] = dep[lst[u]]; } } struct Query {int l, r;} qu[500010]; char t[5000010], pt[1000010]; vector<int> seg[5000010]; vector<pair<int, int>> qv[5000010]; int mxr[5000010]; set<int> st; void solve() { int n = read(), q = read(); scanf("%s", t + 1); int nt = strlen(t + 1); for (int i = 1; i <= n; i++) { scanf("%s", pt + 1); ci++; int len = strlen(pt + 1); pp[i] = ""; for (int x = len; x >= 1; x--) pp[i] += pt[x]; insertp(pt, idxp[i]); inserts(pt, idxs[i]); } dfs(0); buildp(); builds(); dfsp(0); dfss(0); for (int x = 1; x <= n; x++) { int len = pp[x].size(); g[x].resize(len + 2, 0); int u = 0; for (int i = 1; i <= len; i++) { u = ts[u].son[pp[x][i - 1] - 'a']; g[x][i] = sums[u] + g[x][i - 1]; } } match(t); for (int i = 1; i <= nt; i++) { if (lgt[i] == 0) continue; int start = i - lgt[i] + 1; if (start >= 1) seg[start].push_back(i); } for (int i = 1; i <= q; i++) { qu[i].l = read(); qu[i].r = read(); qv[qu[i].l].push_back({qu[i].r, i}); } for (int l = 1; l <= nt; l++) { for (auto tq : qv[l]) { int r = tq.first, id = tq.second; if (st.empty()) {mxr[id] = -1; continue;} auto p = st.upper_bound(r); if (p == st.begin()) {mxr[id] = -1; continue;} p--; if (*p < l) {mxr[id] = -1; continue;} mxr[id] = *p; } for (int r : seg[l]) st.insert(r); } for (int i = 1; i <= q; i++) { ll ans = 0; int l = qu[i].l, r = qu[i].r; int kr = mxr[i]; if (kr == -1) { printf("%lld ", suc[r] - suc[l - 1]); continue; } ans += suc[r] - suc[kr]; int num = tp[lst[pos[kr]]].idx; ans += g[num][kr - l + 1]; printf("%lld ", ans); } } signed main() { #ifdef FIO string f_name = "xxx"; freopen((f_name + ".in").c_str(), "r", stdin); freopen((f_name + ".out").c_str(), "w", stdout); #endif int REP = 1; while (REP--) solve(); return 0; }
|
ps:写得有点烂,请多多包涵。