WA 最后一个 Sub 70pts 求助
查看原帖
WA 最后一个 Sub 70pts 求助
362750
TernaryTree楼主2023/8/27 20:28
#include <bits/stdc++.h>
#define int long long

using namespace std;

const int maxn = 1e6 + 10;
const int mul = 1145141;
const int mod = 0x3b800001;

int n, m, ans;
char s[maxn], t[maxn];
int hs[maxn], ht[maxn], p[maxn];
int lcp[maxn], lcs[maxn];
unordered_map<int, int> pre, suf;
vector<int> vp[maxn], vs[maxn];
int b[2][maxn];

inline int f(int * h, int l, int r) { 
	return (h[r] - h[l - 1] * p[r - l + 1] % mod + mod) % mod; 
}

inline int lowbit(int x) { return x & -x; }
void add(int id, int p, int x) { while (p <= n) b[id][p] += x, p += lowbit(p); }
int query(int id, int p) { int z = 0; while (p) z += b[id][p], p -= lowbit(p); return z; }

signed main() {
	cin >> (s + 1) >> (t + 1);
	n = strlen(s + 1), m = strlen(t + 1);
	hs[0] = ht[0] = 0, p[0] = 1;
	for (int i = 1; i <= max(n, m) + 1; i++) p[i] = p[i - 1] * mul % mod;
	for (int i = 1; i <= n; i++) hs[i] = (hs[i - 1] * mul + s[i]) % mod;
	for (int i = 1; i <= m; i++) ht[i] = (ht[i - 1] * mul + t[i]) % mod;
	for (int i = 1; i <= m; i++) pre[f(ht, 1, i)] = i, suf[f(ht, i, m)] = i;
	for (int i = 1; i <= n; i++) {
		int l = i, r = n;
		while (l <= r) {
			int mid = l + r >> 1;
			if (!pre.count(f(hs, i, mid))) r = mid - 1;
			else l = mid + 1;
		}
		lcp[i] = l - i;
		vp[lcp[i]].push_back(i);
	}
	for (int i = 1; i <= n; i++) {
		int l = 1, r = i;
		while (l <= r) {
			int mid = l + r >> 1;
			if (!suf.count(f(hs, mid, i))) l = mid + 1;
			else r = mid - 1;
		}
		lcs[i] = i - r;
		vs[lcs[i]].push_back(i);
	}
	int cur = 0;
	for (int i = 1; i <= n; i++) add(0, i, 1);
	for (int i = 1; i <= n; i++) {
		if (lcs[i] == m) {
			add(1, i, 1);
			cur += query(0, i - m);
		}
	}
	for (int i = 1; i < m; i++) {
		int j = m - i;
		for (int k : vp[i - 1]) {
			add(0, k, -1);
			if (k + m > n) continue;
			cur -= query(1, n) - query(1, k + m - 1);
		}
		for (int k : vs[j]) {
			add(1, k, 1);
			if (k - m < 1) continue;
			cur += query(0, k - m);
		}
		ans += cur;
	}
	for (int i = 1; i <= n - m + 1; i++) {
		if (lcp[i] == m) {
			ans += i * (i - 1) / 2;
			ans += (n - i - m + 1) * (n - i - m + 2) / 2;
		}
	}
	cout << ans << endl;
	return 0;
}
2023/8/27 20:28
加载中...