P17214 [ICPC 2017 Nanning R] Banned Patterns

· · 题解

Analysis

提供一个哈希思路。

两个等长字符串能够通过字符的全排列相互转换,当且仅当它们内部字符重复出现的位置模式完全相同

我们可以将字符串转化为规范化序列:

两个字符串同构当且仅当它们的规范化序列完全相同。因此,我们可以预处理所有黑名单模式串的规范化哈希值。在待判定的字符串 S 上滑动长度为 L 的窗口时,当窗口从 i 平移到 i+1 时,规范化序列仅有极少数位置发生变化:

  1. 离开的元素 S_i 在旧窗口内为首次出现,其值为 0,移出无影响。
  2. 新进入的元素 S_{i+L}:根据其上一次出现位置 prev_{i+L} 是否在新窗口内,加上其新的相对距离贡献。

这样即可实现滑动窗口哈希的 \mathcal{O}(1) 快速转移。

Detail

  1. 预处理模式串:计算所有模式串的规范化双哈希值,按长度加哈希值存入结构体数组中并排序去重。
  2. 处理待判定串:对于每个查询字符串 S,先预处理出每个位置字符的前驱位置 prev[i] 与后继位置 nxt[i]
  3. 按长度匹配:按出现的不同模式串长度 L 分组,在 S 上滑动长度为 L 的窗口,利用 std::binary_search 在对应的模式串哈希区间内查找。一旦查找到合法匹配立即标记并 Break 剪枝。

:::::success[Code]

#include <iostream>
#include <algorithm>
#include <cstring>

using ull = unsigned long long;

constexpr int MAXN = 5e3 + 10;
constexpr int MAXM = 1e6 + 10;
constexpr ull base1 = 1e6 + 3, mod1 = 1e9 + 7;
constexpr ull base2 = 1e6 + 33, mod2 = 998244353;

ull pow1[MAXM], pow2[MAXM];

struct Hashing {
    ull h1, h2;
    bool operator==(const Hashing& rhs) const { return h1 == rhs.h1 && h2 == rhs.h2; }
    bool operator<(const Hashing& rhs) const { return h1 != rhs.h1 ? h1 < rhs.h1 : h2 < rhs.h2; }
};

struct String {
    int len; Hashing hsh;
    bool operator<(const String& rhs) const { return len != rhs.len ? len < rhs.len : hsh < rhs.hsh; }
    bool operator==(const String& rhs) const { return len == rhs.len && hsh == rhs.hsh; }
} str[MAXN];

int prev[MAXM], nxt[MAXM];
char p[MAXM], s[MAXM];

inline Hashing build(const char* str_ptr, int len) {
    int lst[26];
    std::memset(lst, -1, sizeof lst);

    ull h1 = 0, h2 = 0;
    for (int i = 0; i < len; i++) {
        int c = str_ptr[i] - 'A';
        int val = (lst[c] != -1) ? (i - lst[c]) : 0;
        lst[c] = i;
        h1 = (h1 * base1 + val) % mod1;
        h2 = (h2 * base2 + val) % mod2;
    }
    return {h1, h2};
}

void solve(int _) {
    int n;
    std::cin >> n;

    for (int i = 0; i < n; i++) {
        std::cin >> p;
        int pl = std::strlen(p);
        str[i] = {pl, build(p, pl)};
    }

    std::sort(str, str + n);
    n = std::unique(str, str + n) - str;

    int m;
    std::cin >> m;

    std::cout << "Case #" << _ << ":";

    while (m--) {
        std::cin >> s;
        int sl = std::strlen(s);
        bool flag = false;
        int lst[26];

        std::memset(lst, -1, sizeof lst);
        for (int i = 0; i < sl; i++) {
            prev[i] = -1;
            nxt[i] = sl + 5;
        }

        for (int i = 0; i < sl; i++) {
            int c = s[i] - 'A';
            if (lst[c] != -1) {
                prev[i] = lst[c];
                nxt[lst[c]] = i;
            }
            lst[c] = i;
        }

        for (int i = 0; i < n; ) {
            int pl = str[i].len;
            int nxti = i;
            while (nxti < n && str[nxti].len == pl) { nxti++; }

            if (pl <= sl) {
                ull h1 = 0, h2 = 0;
                for (int k = 0; k < pl; k++) {
                    int val = (prev[k] >= 0) ? (k - prev[k]) : 0;
                    h1 = (h1 * base1 + val) % mod1;
                    h2 = (h2 * base2 + val) % mod2;
                }

                Hashing cur = {h1, h2};

                for (int j = 0; j + pl <= sl; j++) {
                    String targ = {pl, cur};
                    if (std::binary_search(str + i, str + nxti, targ)) {
                        flag = true;
                        break;
                    }

                    if (j + pl < sl) {
                        ull nxth1 = (cur.h1 * base1) % mod1;
                        if (nxt[j] < j + pl) {
                            int dist = nxt[j] - j;
                            ull term = (dist * pow1[pl - dist]) % mod1;
                            nxth1 = (nxth1 + mod1 - term) % mod1;
                        }
                        if (prev[j + pl] >= j + 1) {
                            int dist = (j + pl) - prev[j + pl];
                            nxth1 = (nxth1 + dist) % mod1;
                        }

                        ull nxth2 = (cur.h2 * base2) % mod2;
                        if (nxt[j] < j + pl) {
                            int dist = nxt[j] - j;
                            ull term = (dist * pow2[pl - dist]) % mod2;
                            nxth2 = (nxth2 + mod2 - term) % mod2;
                        }
                        if (prev[j + pl] >= j + 1) {
                            int dist = (j + pl) - prev[j + pl];
                            nxth2 = (nxth2 + dist) % mod2;
                        }

                        cur = {nxth1, nxth2};
                    }
                }
            }

            if (flag) break;
            i = nxti;
        }

        std::cout << ' ' << (flag ? 'Y' : 'N');
    }
    std::cout << '\n';
}

int main() {
    std::ios::sync_with_stdio(false);
    std::cin.tie(nullptr);

    pow1[0] = pow2[0] = 1;
    for (int i = 1; i <= 1000000; i++) {
        pow1[i] = (pow1[i - 1] * base1) % mod1;
        pow2[i] = (pow2[i - 1] * base2) % mod2;
    }

    int T; std::cin >> T;
    for (int _ = 1; _ <= T; _++){ solve(_); }

    return 0;
}

:::::