题解:SP1676 GEN - Text Generator

· · 题解

:::info[提示] 本文使用了 Deepseek 进行了润色,严格保证人类贡献大于 AI 贡献。 :::

MX 题单里做到的。

前置芝士:AC 自动机,矩阵快速幂。

Problem.

给定 N 个字符串和 L,求长度为 L 的、包含给定的至少一个字符串的数量,结果对 10007 取模(给定的字符串全部由大写字母组成)。

Solution.

一眼 dp。

直接算“包含”显然不好算,正难则反,不妨从反向入手。

显然的,答案 = 所有可能的字符串总数 - 不包含任何给定模式串的字符串的数量

注意到总字符串的数量为 26^L

因此,我们可以计算出长度为 L 且不包含任何给定字符串的字符串的数量,问题就解决了。

构建 AC 自动机,把所有模式串插入 Trie 树并构建 fail 指针。

先讲如何处理非法状态。

显然有以下两种情况不合法:

将它们标记出来,后面方便处理。

然后考虑怎么 dp。

定义 dp_{i,j} 表示已经生成了长度为 i 的字符串,且当前在 AC 自动机的节点 j 上,并且这个字符串到目前为止没有包含任何模式串的方案数。

初始化显然 dp_{0,0}=1

定义 \text{bad} 为不合法字符串的数量,\text{bad} = \sum_{j \in \text{Safe}} dp_{L,j},其中 \text{Safe} 表示所有合法节点的集合。

考虑如何状态转移。

设当前在节点 u 且当前字符串合法。

枚举下一个节点 c,在 AC 自动机上走到 v=tr_{u,c}

如果 v 不合法直接跳过。

否则就继续走下去。

则状态转移方程为 dp_{i+1,v}=dp_{i+1,v}+dp_{i,u}

形式化一点:

dp_{i+1,v}=\sum_{u \in Safe} dp_{i,u} \times [tr_{u,c}=v]

其中 i0L-1,共递推 L 步。

注意到 L 巨大,直接算显然时间会爆,因为是线性递推,考虑用矩阵快速幂优化。

矩阵快速幂的实现就简单点。

定义状态向量 F_i 表示长度为 i 时各个节点的方案数。

构造转移矩阵 MM_{u,v} 表示从节点 u 一步到达节点 v 的方案数(只保留安全节点的转移)。

F_{i+1} = F_i \times M,所以 F_L = F_0 \times M^L

用矩阵快速幂计算 M^L,用 26^L 减完得到答案。

然后这题就做完了。

时间复杂度 \Omicron(S^3 \times \log L),空间为 \Omicron(S^2)S 表示合法节点数。

贴个代码。

:::info[代码]

#include<bits/stdc++.h>
#define BUF 1<<20
#define IL inline
#define ll long long
#define ri register int
#define F(i,a,b) for(ri i=a;i<=b;i++)
#define FF(i,a,b) for(ri i=b;i>=a;i--)
#define u64 uint64_t
#define ull unsigned long long
#define i128 __int128
#define vec vector
#define vi vector<int>
#define vll vector<ll>
#define vb vector<bool>
#define prq priority_queue
#define pii pair<int,int>
#define pill pair<int,ll>
#define plli pair<ll,int>
#define um unordered_map
#define mii map<int,int>
#define us unordered_set
#define pb(x) push_back(x)
#define fi first
#define se second
#define fr() front()
#define bk() back()
#define beg() begin()
#define Fill(a,b) memset(a,b,sizeof(a))
using namespace std;
const int N=105,mod=1e4+7;
const ll inf=0x3f3f3f3f3f3f3f3fLL;
const double eps=1e-9;
char buf[BUF],*p1=buf,*p2=buf;
#define getchar_unlocked()((p1==p2)&&(p2=(p1=buf)+fread(buf,1,BUF,stdin),p1==p2)?EOF:*p1++)
int n,l,tot,ans,bad;
string s;
struct Matrix{
    int a[N][N],n;
    Matrix(int n=0,bool id=0):n(n){
        Fill(a,0);
        if(id){
            F(i,0,n-1) a[i][i]=1;
        }
    }
    Matrix operator *(const Matrix&oth)const{
        Matrix res(n);
        F(i,0,n-1){
            F(k,0,n-1){
                if(!a[i][k]) continue;
                F(j,0,n-1){
                    res.a[i][j]=(res.a[i][j]+a[i][k]*oth.a[k][j])%mod;
                }
            }
        }
        return res;
    }
};
IL Matrix qpow1(Matrix b,int e){
    Matrix res(b.n,1);
    while(e){
        if(e&1) res=res*b;
        b=b*b;
        e>>=1;
    }
    return res;
}
struct ACAM{
    int tr[N][26],fail[N],idx;
    bool dan[N];
    IL void init(){
        Fill(tr,0),Fill(fail,0),Fill(dan,0),idx=0;
    }
    IL void insert(string s){
        int p=0;
        for(char ch:s){
            int c=ch-'A';
            if(!tr[p][c]) tr[p][c]=++idx;
            p=tr[p][c];
        }
        dan[p]=1;
    }
    IL void build(){
        queue<int> q;
        F(i,0,25){
            if(tr[0][i]){
                q.push(tr[0][i]);
            }
        }
        while(!q.empty()){
            int u=q.fr(); q.pop();
            dan[u]|=dan[fail[u]];
            F(i,0,25){
                if(tr[u][i]){
                    fail[tr[u][i]]=tr[fail[u]][i];
                    q.push(tr[u][i]);
                }
                else tr[u][i]=tr[fail[u]][i];
            }
        }
    }
    IL Matrix buildM(){
        Matrix mat(idx+1);
        F(u,0,idx){
            if(dan[u]) continue;
            F(i,0,25){
                if(!dan[tr[u][i]]) mat.a[u][tr[u][i]]++;
            }
        }
        return mat;
    }
} acam;
IL int qpow2(int e){
    int a=26,res=1;
    while(e){
        if(e&1) res=res*a%mod;
        a=a*a%mod;
        e>>=1;
    }
    return res;
}
IL int read(){
    int k=0,f=1;
    char c=getchar_unlocked();
    while(c<'0'||c>'9'){
        if(c=='-') f=-1;
        c=getchar_unlocked();
    }
    while(c>='0'&&c<='9') k=k*10+c-'0',c=getchar_unlocked();
    return k*f;
}
IL void write(int x){
    if(x<0) putchar('-'),x=-x;
    if(x<10) putchar(x+'0');
    else write(x/10),putchar(x%10+'0');
}

int main(){
    ios::sync_with_stdio(false);
    cin.tie(0);
//  freopen(".in","r",stdin);
//  freopen(".out","w",stdout);
    while(cin>>n>>l) {
        bad=0;
        acam.init();
        F(i,1,n){
            cin>>s;
            acam.insert(s);
        }
        acam.build();
        Matrix t=acam.buildM();
        Matrix fi=qpow1(t,l);
        F(i,0,acam.idx){
            bad=(bad+fi.a[0][i])%mod;
        }
        tot=qpow2(l);
        ans=(tot-bad+mod)%mod;
        write(ans);
        putchar('\n');
    }
    return 0;
}

:::

这题有个弱化版 P4052,但是那题 N 更大 L 更小,线性算即可。