peiwenjun's blog 没有知识的荒原

CF1801G A task for substrings 题解

题目描述

给定文本串 \(s\) 和若干模式串 \(t\)\(q\) 次询问 \(s[l\sim r]\) 有多少子串在 \(t\) 中出现。

定义两个子串不同,当且仅当它们在 \(s\) 中的位置不同。

数据范围

  • \(1\le |s|\le 5\cdot 10^6,1\le\sum|t|\le 10^6,1\le q\le 5\cdot 10^5\)

时间限制 \(\texttt{4s}\) ,空间限制 \(\texttt{1GB}\)

分析

多模匹配,先建 \(\text{AC}\) 自动机。

\(AC\) 自动机不好处理子串信息,但是处理前缀是容易的,统计一下每个点 fail 树跳到根的权值和即可。

考虑容斥,记 \(sum_i\)\(s[1\sim i]\) 的匹配次数。

那么对于每组询问 \([l,r]\) ,答案为 \(sum_r-sum_{l-1}\) ,再减去左端点 \(\in[1,l-1]\) ,右端点 \(\in[l,r]\) 的贡献。

接下来考虑如何计算跨过 \([l-1,l]\) 的贡献,先将所有模式串的反串也建一个 \(\text{AC}\) 自动机。

对每个模式串 \(t\) ,在正串的 \(\text{ACAM}\) 上跑一个前缀 \(t[1\sim j]\) ,在反串的 \(\text{ACAM}\) 上跑一个后缀 \(t[j+1\sim|t|]\) ,将这两个节点连边。

那么对于每个跨过 \([l-1,l]\) 的模式串,从跨越位置切开,它会唯一对应一条边。

记正串的 \(\text{ACAM}\) 上跑 \(s[1\sim l-1]\) 会到达节点 \(x\) ,反串的 \(\text{ACAM}\) 上跑 \(s[l\sim r]\) 会到达节点 \(y\) ,答案为\(x\) 的祖先中任选一个点、在 \(y\) 的祖先中任选一个点,两点之间的边数

首当其冲的问题是 \(x,y\) 怎么求。

注意到如果 \(s_1\) 是 \(s_2\) 的后缀,那么 \(s1\) 到达的节点一定是 \(s_2\) 到达的节点在 fail 树上的祖先。

因此记录一下每个点的字符串长度( trie 树上的深度),在树上倍增跳即可。

再将询问离线并挂在 \(x\) 上, \(dfs\) 第一棵 fail 树,加入当前节点在第二棵树上的边,回溯时撤销即可。

询问就是在第二棵 fail 树上点到根的路径求和。

单点加、链求和显然可以转化成子树加、单点求值,树状数组维护即可。

时间复杂度 \(\mathcal O(|s|+(\sum|t|+q)\log\sum|t|)\)

#include<bits/stdc++.h>
#define ll long long
#define fi first
#define se second
#define mp make_pair
#define pii pair<int,int>
using namespace std;
const int maxn=1e6+5,maxm=5e6+5;
int l,m,n,q,r;
ll c[maxn],res[maxn],sum[maxm];
int tmp[2][maxn];
char s[maxm],ch[maxn];
string t[maxn];
vector<int> h[maxn];
vector<pii> vec[maxn];
struct acam
{
    int op,cnt,tot;
    int sz[maxn],dep[maxn],dfn[maxn],val[maxn];
    int ch[maxn][26],fail[maxn];
    int pos[maxm];
    vector<int> g[maxn],fa[maxn];
    void insert(string s)
    {
        if(op) reverse(s.begin(),s.end());
        int n=s.size(),p=0;
        for(int i=0;i<n;i++)
        {
            int k=s[i]-'a';
            if(!ch[p][k]) ch[p][k]=++tot,dep[tot]=dep[p]+1;
            p=ch[p][k];
        }
        val[p]++;
    }
    void getfail()
    {
        queue<int> q;
        for(int i=0;i<=25;i++) if(ch[0][i]) q.push(ch[0][i]);
        while(!q.empty())
        {
            int u=q.front();
            q.pop();
            g[fail[u]].push_back(u),val[u]+=val[fail[u]];
            for(int i=0;i<=25;i++)
                if(ch[u][i]) fail[ch[u][i]]=ch[fail[u]][i],q.push(ch[u][i]);
                else ch[u][i]=ch[fail[u]][i];
        }
    }
    void dfs(int u)
    {
        sz[u]=1,dfn[u]=++cnt;
        for(auto v:g[u])
        {
            if(op)
            {
                fa[v].resize(20),fa[v][0]=u;
                for(int i=1;i<=19;i++) fa[v][i]=fa[fa[v][i-1]][i-1];
            }
            dfs(v),sz[u]+=sz[v];
        }
    }
    int get(int u,int len)
    {
        if(dep[u]<=len) return u;
        for(int i=19;i>=0;i--) if(dep[fa[u][i]]>len) u=fa[u][i];
        return fa[u][0];
    }
}t1,t2;
void add(int x,int v)
{
    while(x<=t2.cnt) c[x]+=v,x+=x&(-x);
}
int query(int x)
{
    int res=0;
    while(x) res+=c[x],x-=x&(-x);
    return res;
}
void dfs(int u)
{
    for(auto v:h[u]) add(t2.dfn[v],1),add(t2.dfn[v]+t2.sz[v],-1);
    for(auto p:vec[u]) res[p.se]-=query(t2.dfn[p.fi]);
    for(auto v:t1.g[u]) dfs(v);
    for(auto v:h[u]) add(t2.dfn[v],-1),add(t2.dfn[v]+t2.sz[v],1);
}
int main()
{
    scanf("%d%d%s",&n,&q,s+1),m=strlen(s+1),t2.op=1;
    for(int i=1;i<=n;i++)
    {
        scanf("%s",ch),t[i]=ch;
        t1.insert(t[i]),t2.insert(t[i]);
    }
    t1.getfail(),t2.getfail();
    for(int i=1,p=0;i<=m;i++) t1.pos[i]=p=t1.ch[p][s[i]-'a'],sum[i]=sum[i-1]+t1.val[p];
    for(int i=m,p=0;i>=1;i--) t2.pos[i]=p=t2.ch[p][s[i]-'a'];
    t1.dfs(0),t2.fa[0].resize(20),t2.dfs(0);
    for(int i=1;i<=n;i++)
    {
        int l=t[i].size();
        for(int j=0,p=0;j<l;j++) tmp[0][j]=p=t1.ch[p][t[i][j]-'a'];
        for(int j=0,p=0;j<l;j++) tmp[1][j]=p=t2.ch[p][t[i][l-1-j]-'a'];
        for(int j=0;j<=l-2;j++) h[tmp[0][j]].push_back(tmp[1][l-2-j]);
    }
    for(int i=1;i<=q;i++)
    {
        scanf("%d%d",&l,&r),res[i]=sum[r]-sum[l-1];
        int x=t1.pos[l-1],y=t2.get(t2.pos[l],r-l+1);
        vec[x].push_back(mp(y,i));
    }
    dfs(0);
    for(int i=1;i<=q;i++) printf("%lld ",res[i]);
    putchar('\n');
    return 0;
}

posted on 2023-04-20 23:29  peiwenjun  阅读(10)  评论(0)    收藏  举报

导航