SAM 学习笔记 | #2139. 字符串
SAM 是一种自动机,可以在 \(\mathcal O(n)\) 的时间内构建,并以 \(O(n)\) 的时间维护子串相关信息。
SAM 的结构
SAM 分为两部分,一部分是自动机,另一部分是由后缀链接 link 构成的内向树 parentTree。
自动机上,接受原串的任意后缀,并最终抵达终止节点。
考虑到每个子串是原串一个后缀的一个前缀,以此可以在 SAM 上高效维护子串相关信息。
SAM 中每个节点的信息是高度压缩的。每个节点实际上维护一个 endpos 等价类。
每个等价类中较短的串都是较长的串的后缀。且串长取遍 \([\text{len}(\text{link}(u))+1,\text{len}(u)]\)。
对于每个等价类,我们维护 link,指向其最长的 endpos 不同的后缀。
最终会得到一个以 0 为根的内向树,满足父亲的 endpos 为其所有儿子的 endpos 的并,我们称其为 parentTree。
通过在 parentTree 上处理树上问题可以维护绝大多数的子串信息。
建 SAM
考虑增量法,每次添加下一个字符。
维护一个 \(last\) 表示当前整个字符串在自动机上的编号。
加入新字符 \(t\)。取 \(cur\) 表示新的节点。
在 parentTree 上跳 \(last\) 的祖先,记当前枚举到 \(p\)。
若 \(p\) 没有值为 \(t\) 的出边,则连一条 \(p\) 到 \(cur\) 的值为 \(t\) 的出边。
否则当枚举到一个 \(p\) 存在连到 \(q\) 的值为 \(t\) 的出边。
若 \(\text{len}(q)=\text{len}(p)+1\),则令 \(\text{link}(cur)=q\)。
否则说明需要拆分等价类。(对应 "xabyab" 插入 y 时的情形)
新建一个节点 copy,继承 \(q\) 的 link 和出边信息。
同时令 \(\text{len}(copy)=\text{len}(p)+1\),\(\text{link}(q)=\text{link}(cur)=copy\)。
接下来,枚举 \(p\) 在 parentTree 上的祖先 \(p'\),将 \(p'\) 连向 \(q\) 的边改向 \(copy\),直到 \(p'\) 值为 \(t\) 的出边不指向 \(q\) 则退出。
观察这个过程可以发现其空间和时间均为线性的。
#2139 字符串 题解
考虑题目中的条件即为 endpos 中至少有 \(k-1\) 个长度为 \(\text{len}+1\) 的 gap。
那么我们在 parentTree 上启发式合并,维护 endpos 集合。同时另外用一个线段树维护所有 gap 的长度。
枚举到一个节点时通过线段树二分确定需要的 len 长度,其属于一区间,若其与该节点的 len 区间有交,则存在解。
若我们需要最小化字典序,考虑我们是在 parentTree 上操作,则可以将原串反转。
枚举到一个节点的时候按儿子倒数第 \(\text{len}(u)+1\) 项的值排序,即可做到最小化字典序,这个过程可以直接取任一 endpos 同时在原串上查询。
#include <algorithm>
#include <iostream>
#include <assert.h>
#include <string.h>
#include <vector>
#include <set>
typedef long long i64;
const int N = 1e5 + 7, M = 26;
struct node {
int len, link, cnt;
int to[M];
} tr[N<<1];
int n, last, idx;
int pos[N<<1];
inline void init() {
for(int i = 0; i < idx; ++i) pos[i] = 0;
last = 0, idx = 1;
memset(tr + last, 0, sizeof(node));
tr[last].link = -1;
n = 0;
}
inline void append(int c) {
int cur = idx++, p = last;
memset(tr + cur, 0, sizeof(node));
tr[cur].len = tr[p].len + 1;
for(; ~p && !tr[p].to[c]; p = tr[p].link)
tr[p].to[c] = cur;
if(p == -1) tr[cur].link = 0;
else {
int q = tr[p].to[c];
if(tr[q].len == tr[p].len + 1) tr[cur].link = q;
else {
int copy = idx++; memcpy(tr + copy, tr + q, sizeof(node));
tr[copy].len = tr[p].len + 1, tr[copy].cnt = 0;
tr[q].link = tr[cur].link = copy;
for(; ~p && tr[p].to[c] == q; p = tr[p].link)
tr[p].to[c] = copy;
}
}
++tr[cur].cnt, last = cur, pos[cur] = ++n;
}
struct SegT {
struct node {
int ls, rs, sum;
} tr[N*64];
int idx = 1;
inline int newnode() {
return ++idx, tr[idx].ls = tr[idx].rs = tr[idx].sum = 0, idx;
}
void clear() {
idx = 0;
}
#define self tr[rt]
void update(int& rt, int x, int y, int l=1, int r=n) {
if(!rt) rt = newnode();
int m = (l + r) >> 1;
self.sum += y;
if(l < r) {
if(x <= m) update(self.ls, x, y, l, m);
else update(self.rs, x, y, m+1, r);
}
}
int queryR(int& rt, int k, int l=1, int r=n) {
if(!rt) return k ? -1 : r;
if(l == r) return self.sum == k ? l : -1; // exactly k
int m = (l + r) >> 1;
if(k <= tr[self.rs].sum) return queryR(self.rs, k, m+1, r);
else return queryR(self.ls, k - tr[self.rs].sum, l, m);
}
int queryL(int& rt, int k, int l=1, int r=n) {
if(!rt) return l;
if(l == r) return l;
int m = (l + r) >> 1;
if(k < tr[self.rs].sum) return queryL(self.rs, k, m+1, r);
else return queryL(self.ls, k - tr[self.rs].sum, l, m);
}
} T;
int rt[N*2], id[N*2];
std::set<int> set[N*2]; // maintain endpos
std::basic_string<int> g[2*N];
inline int merge(int x, int y) {
if(set[x].size() > set[y].size()) std::swap(x, y);
auto &a = set[x], &b = set[y];
for(auto& p: a) {
auto [it, res] = b.insert(p);
int pr = -1, ne = -1;
if(it != b.begin()) { auto _it = it; pr = *--_it; }
if(++it != b.end()) ne = *it;
if(~pr) T.update(rt[y], p - pr, 1);
if(~ne) T.update(rt[y], ne - p, 1);
if(~pr && ~ne) T.update(rt[y], ne - pr, -1);
}
set[x].clear();
return y;
}
inline void solve() {
std::string str; int k;
std::cin >> str >> k; --k;
std::reverse(str.begin(), str.end());
init(); for(auto& c: str) append(c-'a');
for(int i = 0; i < idx; ++i) {
g[i].clear(), set[i].clear(), id[i] = rt[i] = 0;
}
for(int i = 1; i < idx; ++i) g[tr[i].link] += i;
str = ' ' + str;
T.clear();
static std::pair<int, int> ans[N*2];
for(int i = 1; i <= idx; ++i) ans[i] = {0, 0};
[](auto&&f, auto&&...args) { f(f, args...); } (
[&](auto&& ptr, int u) -> void {
id[u] = u;
if(pos[u]) {
set[u].insert(pos[u]);
}
std::vector<std::pair<char, int>> sons; // to sort
for(int& v: g[u]) {
ptr(ptr, v);
int w = *set[id[v]].begin()-tr[u].len;
sons.emplace_back(str[w], v);
id[u] = merge(id[u], id[v]);
}
// lets check if we can obtain an answer
int lenL = T.queryL(rt[id[u]], k), lenR = T.queryR(rt[id[u]], k)-1;
int myL = tr[tr[u].link].len + 1, myR = tr[u].len;
if(std::min(myR, lenR) >= std::max(myL, lenL)) {
// omg here is indeed an answer! any of the endpos is valid.
int len = std::max(myL, lenL), p = *set[id[u]].begin();
ans[u] = {p - len + 1, p};
}
std::sort(sons.begin(), sons.end());
g[u].clear(); for(auto& [x, y]: sons) g[u] += y;
}, 0
);
int flag = 0;
try{ [](auto&& f, auto&&...args) { f(f, args...); } (
[&](auto&& ptr, int u) -> void {
if(ans[u].first) {
// its non-zero! lets bet it must be the answer.
auto tmp = str.substr(ans[u].first, ans[u].second - ans[u].first + 1);
std::reverse(tmp.begin(), tmp.end());
std::cout << tmp << "\n";
throw 520;
}
for(int& v: g[u]) {
ptr(ptr, v);
}
}, 0
); } catch(int anything) {
flag = 1;
}
if(!flag) {
std::cout << "-1\n";
}
}
int main() {
std::ios::sync_with_stdio(0), std::cin.tie(0), std::cout.tie(0);
int t; std::cin >> t; while(t--) solve();
}
本文来自博客园,作者:CuteNess,转载请注明原文链接:https://www.cnblogs.com/CuteNess/p/21895275

浙公网安备 33010602011771号