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();
}
posted @ 2026-07-25 00:00  CuteNess  阅读(12)  评论(0)    收藏  举报