字符串
双模字符串哈希
// 双哈希。get_s(l, r) 使用 0 下标闭区间,equal/lcp 使用的也是 0 下标。
// 哈希只能把冲突概率降得很低,不能从数学上保证绝对无冲突。
// 如果题目区间是 1 下标,推荐调用 get_s(l - 1, r - 1)。
// 也可以先写 s = ' ' + s,此时原字符位于 1..原长度,直接传题目的 l、r 也能查询;
// 但 size() 会比原长度多 1,使用 lcp() 时也要记得这个占位字符属于字符串的一部分。
struct DoubleHash
{
const int P = 13331;
const int MOD1 = 998244353;
const int MOD2 = 1000000007;
int n = 0;
vector<int> p1, p2, hs1, hs2;
DoubleHash()
{
}
DoubleHash(const string &s)
{
init(s);
}
void init(const string &s)
{
n = s.size();
p1.assign(n + 1, 0);
p2.assign(n + 1, 0);
hs1.assign(n + 1, 0);
hs2.assign(n + 1, 0);
p1[0] = p2[0] = 1;
for (int i = 1; i <= n; i++)
{
p1[i] = p1[i - 1] * P % MOD1;
p2[i] = p2[i - 1] * P % MOD2;
int ch = (unsigned char)s[i - 1] + 1;
hs1[i] = (hs1[i - 1] * P + ch) % MOD1;
hs2[i] = (hs2[i - 1] * P + ch) % MOD2;
}
}
pair<int, int> get_s(int l, int r) const
{
if (l > r)
{
return {0, 0};
}
int x1 = (hs1[r + 1] - hs1[l] * p1[r - l + 1] % MOD1 + MOD1) % MOD1;
int x2 = (hs2[r + 1] - hs2[l] * p2[r - l + 1] % MOD2 + MOD2) % MOD2;
return {x1, x2};
}
bool equal(int l1, int r1, int l2, int r2) const
{
if (r1 - l1 != r2 - l2)
{
return false;
}
return get_s(l1, r1) == get_s(l2, r2);
}
// 返回两个后缀从 l1、l2 开始的最长公共前缀长度,最多比较 limit 个字符。
int lcp(int l1, int l2, int limit) const
{
limit = min({limit, n - l1, n - l2});
int l = 0, r = limit;
while (l < r)
{
int mid = (l + r + 1) >> 1;
if (get_s(l1, l1 + mid - 1) == get_s(l2, l2 + mid - 1))
{
l = mid;
}
else
{
r = mid - 1;
}
}
return l;
}
int size() const
{
return n;
}
};
字符 Trie
// ============================================================
// 小写字母 Trie
// ============================================================
// end[p] :恰好在节点 p 结束的字符串数量。
// pass[p]:经过节点 p 的字符串数量,因此可用于查询前缀数量。
// erase 只减少计数,不回收节点;竞赛中通常更稳,也便于重复插入/删除。
// 如果字符集不是 a-z,修改 SIGMA 和 getid(),并同步修改 array 的大小。
// ============================================================
class Trie
{
private:
static constexpr int SIGMA = 26;
vector<array<int, SIGMA>> tr;
vector<int> pass, ending;
int getid(char c)
{
return c - 'a';
}
public:
Trie()
{
clear();
}
void clear()
{
tr.assign(1, array<int, SIGMA>{});
pass.assign(1, 0);
ending.assign(1, 0);
}
void insert(const string &s)
{
int p = 0;
pass[p]++;
for (char c : s)
{
int u = getid(c);
if (!tr[p][u])
{
tr[p][u] = tr.size();
tr.push_back(array<int, SIGMA>{});
pass.push_back(0);
ending.push_back(0);
}
p = tr[p][u];
pass[p]++;
}
ending[p]++;
}
int query(const string &s)
{
int p = 0;
for (char c : s)
{
int u = getid(c);
if (!tr[p][u])
{
return 0;
}
p = tr[p][u];
}
return ending[p];
}
int count_prefix(const string &prefix)
{
int p = 0;
for (char c : prefix)
{
int u = getid(c);
if (!tr[p][u])
{
return 0;
}
p = tr[p][u];
}
return pass[p];
}
bool starts_with(const string &prefix)
{
return count_prefix(prefix) > 0;
}
bool erase(const string &s)
{
if (!query(s))
{
return false;
}
int p = 0;
pass[p]--;
for (char c : s)
{
p = tr[p][getid(c)];
pass[p]--;
}
ending[p]--;
return true;
}
int count_all()
{
return pass[0];
}
int nodes()
{
return tr.size();
}
};
Z 函数
// z[i]:s 与 s[i..] 的最长公共前缀长度。
vector<int> z_function(const string& s) {
int n = (int)s.size();
vector<int> z(n);
z[0] = n;
for (int i = 1, l = 0, r = 0; i < n; ++i) {
if (i <= r) z[i] = min(r - i + 1, z[i - l]);
while (i + z[i] < n && s[z[i]] == s[i + z[i]]) ++z[i];
if (i + z[i] - 1 > r) {
l = i;
r = i + z[i] - 1;
}
}
return z;
}
// pattern 在 text 中的全部出现位置(0 下标)。
vector<int> find_occurrences(string pattern, const string& text) {
string s = pattern + "#" + text;
auto z = z_function(s);
vector<int> pos;
for (int i = (int)pattern.size() + 1; i < (int)s.size(); ++i)
if (z[i] >= (int)pattern.size())
pos.push_back(i - (int)pattern.size() - 1);
return pos;
}