CF528D 思路分享(FFT/NTT)
https://codeforces.com/contest/528/problem/D
题意
给定标准串 \(s\) 和模式串 \(t\),仅由 \(ACGT\) 四种字符组成.
定义模糊匹配:若 \(s[i-k,\cdots, i+k]\) 存在 \(t_j\),则 \(s_i\) 与 \(t_j\) 匹配.
求 \(s\) 中能够模糊匹配成功的起点数量.
\(1\le |t| \le |s| \le 2\cdot 10^5\),\(0\le k \le 2\cdot 10^5\).
思路
由于字符集数量很少,考虑枚举字符 \(c\).
记 \(|S|=n\),\(|T|=m\).
定义 \(A_i\) 为 \(s_i\) 能否模糊匹配 \(c\),这可以通过前缀和处理.
定义 \(B_j\) 为 \(t_j\) 是否为字符 \(c\).
当 \(B_j=1\) 时,\(A_{i+j}\) 必须为 \(1\),否则起点 \(i\) 不合法,即
\[\sum_{j=0}^{m-1}{B_j\cdot(1-A_{i+j})}=0
\]
将 \(A\) 取反,反转 \(B\)
\[\sum_{j=0}^{m-1}{B_{m-1-j}\cdot A_{i+j}}=0
\]
注意到下标之和为 \(m-1+i\) ,这是卷积的形式,使用 \(FFT/NTT\) 计算即可,累加多个卷积的结果,最后为 \(0\) 就是合法的起点.
时间复杂度 \(\mathcal{O}(4L\log L)\),\(L\) 是参与卷积计算的长度.
代码
//author:kzssCCC
#include <bits/stdc++.h>
using namespace std;
using ll = long long;
const int MOD = 998244353;
ll qpow(ll a,ll b){
ll res = 1;
while (b){
if (b&1){
res = res*a%MOD;
}
a = a*a%MOD;
b >>= 1;
}
return res;
}
const int g = 3;
const int invg = qpow(g,MOD-2);
void ntt(vector<ll>& vec,int n,int op){
int B = 31-__builtin_clz(n);
for (int i=0;i<n;i++){
int rev = 0;
for (int j=0;j<B;j++){
rev <<= 1;
rev |= i>>j&1;
}
if (i<rev){
swap(vec[i],vec[rev]);
}
}
for (int len=2;len<=n;len<<=1){
ll w = qpow(op==0?g:invg,(MOD-1)/len);
for (int i=0;i<n;i+=len){
ll cur = 1;
for (int j=0;j<len/2;j++){
ll u = vec[i+j];
ll v = vec[i+j+len/2];
vec[i+j] = (u+cur*v%MOD)%MOD;
vec[i+j+len/2] = (u-cur*v%MOD+MOD)%MOD;
cur = cur*w%MOD;
}
}
}
if (op==1){
for (int i=0;i<n;i++){
vec[i] = vec[i]*qpow(n,MOD-2)%MOD;
}
}
}
void solve(){
int n,m,K;
cin >> n >> m >> K;
string s,t;
cin >> s >> t;
int N = 1;
while (N<n+m-1){
N<<=1;
}
vector<ll> P(N);
for (auto ch:"ACGT"){
vector<ll> pre(n+1);
for (int i=1;i<=n;i++){
pre[i] = pre[i-1]+(s[i-1]==ch);
}
vector<ll> a(n);
for (int i=0;i<n;i++){
int l = max(0,i-K);
int r = min(n-1,i+K);
a[i] = !(pre[r+1]-pre[l]);
}
vector<ll> b(m);
for (int i=0;i<m;i++){
b[i] = t[i]==ch;
}
reverse(b.begin(),b.end());
a.resize(N);
b.resize(N);
ntt(a,N,0);
ntt(b,N,0);
vector<ll> c(N);
for (int i=0;i<N;i++){
c[i] = a[i]*b[i]%MOD;
}
ntt(c,N,1);
for (int i=0;i<N;i++){
P[i] = (P[i]+c[i])%MOD;
}
}
int res = 0;
for (int i=m-1;i<=n-1;i++){
if (P[i]==0){
res++;
}
}
cout << res << '\n';
}
int main(){
ios::sync_with_stdio(false);
cin.tie(0);
int t = 1;
// cin >> t;
while (t--) solve();
return 0;
}

浙公网安备 33010602011771号