数位dp学习笔记

· · 算法·理论

Part1

一种最基础的数位dp,例如 windy数 此类的数位dp入门题,一般我写的都是记忆化搜索,没怎么写过递推。

将 P2602 [ZJOI2010] 数字计数 作为例题来讲解,从高位往低位 dp,我们分开来求每个数出现的次数,首先设计状态 f[N][N] 表示还剩下 i 位前面已经有了 jnow

放个模板: ::::info[code]

#include<bits/stdc++.h>
#define ll long long
using namespace std;
const ll mod=1e9+7;
ll T,n,l,r,c[20],f[20][20],ans,now;
ll dp(int dep,int sta,int lim,int qd0){
    if(qd0==0) sta=0;
    if(dep==0) return sta;
    if(!lim && qd0 && f[dep][sta]!=-1) return f[dep][sta];
    int up=lim?c[dep]:9;
    ll res=0;
    for(int i=0;i<=up;++i){
        res+=dp(dep-1,sta+(i==now),lim&&(i==up),qd0||(i!=0));
    }
    if(!lim && qd0) f[dep][sta]=res;
    return res;
}
ll calc(ll x){
    int cnt=0;
    while(x){c[++cnt]=x%10;x/=10;}
    return dp(cnt,0,1,0);
}
signed main(){
    ios::sync_with_stdio(0),cin.tie(0),cout.tie(0);
    cin>>l>>r;
    for(int i=0;i<=9;i++){
        now=i;memset(f,-1,sizeof f);
        cout<<calc(r)-calc(l-1)<<" ";
    }
    return 0;
}

::::

Part 2

例如 P2106 Sam 数 当数位达到 10^{18} 的时候就很自然的想到用矩阵快速幂来优化数位dp。 ::::info[code]

#include<bits/stdc++.h>
#define ll long long
using namespace std;
const int mod=1e9+7;
ll k,cnt;
struct matrix{
    ll a[15][15],nn,mm;
    matrix(ll nn=0,ll mm=0):nn{nn},mm{mm}{memset(a,0,sizeof a);}
    matrix friend operator*(const matrix &x,const matrix &y){
        matrix z(x.nn,y.mm);
        for(int i=1;i<=x.nn;i++)
            for(int j=1;j<=y.mm;j++)
                for(int k=1;k<=x.mm;k++)
                    z.a[i][j]=(z.a[i][j]+x.a[i][k]*y.a[k][j])%mod;
        return z;
    }
    matrix friend operator^(matrix &x,const long long&y){
        matrix z(x.nn,x.mm),tmp=x;
        for(int i=1;i<=x.nn;i++) z.a[i][i]=1;
        for(ll i=y;i;i>>=1){
            if(i&1) z=z*tmp;
            tmp=tmp*tmp;
        }
        return z;
    }
}x(10,10),ans(1,10);
int main(){
    ios::sync_with_stdio(0),cout.tie(0),cin.tie(0);
    cin>>k;
    for(int i=1;i<=10;i++){
        for(int j=max(1,i-2);j<=min(10,i+2);j++){
            x.a[i][j]=1;
        }
    }
    if(k==1){cout<<10;return 0;}
    for(int i=2;i<=10;i++) ans.a[1][i]=1;
    x=x^(k-1);
    ans=ans*x;
    for(int i=1;i<=10;i++) cnt=(cnt+ans.a[1][i])%mod;
    cout<<cnt;
    return 0;
}

::::

Part 3

当我们需要求的值并不是个数,而是满足条件的数的平方和或 \sum if(i) 这种类型时,我们需要记录多个变量来维护。

例题

link,我们不管题目中所需要的条件,因为这是 Part 1 的事情,那么我们考虑如何维护平方和那么我们就要维护三个值,个数,总和,平方和。

前两个值好维护,令 node tmp=dp(dep-1,lim&&(i==up)) 那么根据 (a+b)^2=a^2+b^2+2ab 可知

res.cnt+=tmp.cnt;
res.sum+=tmp.cnt*i*ten[dep-1]+tmp.sum;
res.ans+=tmp.ans+i*i*ten[dep-1]*ten[dep-1]*tmp.cnt+2*i*ten[dep-1]*tmp.sum;

::::info[code]

#include<bits/stdc++.h>
using namespace std;
const long long mod=1e9+7;
long long n,a,b,ans,c[22],ten[18];
struct node{
    long long f,g,h;
    node(){
        f=0,g=0,h=0;
    }
    node(long long _f,long long _g,long long _h):f(_f),g(_g),h(_h){}
}f[22][8][8];
node dp(int dep,int sta1,int sta2,int lim){
    if(dep==0) return {sta1!=0 && sta2!=0,0,0};
    if(lim==0 && f[dep][sta1][sta2].f!=-1) return f[dep][sta1][sta2];
    int up=lim?c[dep]:9;
    node ret{0,0,0};
    for(int i=0;i<=up;i++){
        if(i!=7){
            node tmp=dp(dep-1,(sta1+i)%7,(sta2*10+i)%7,lim&&(i==up));
            ret.f+=tmp.f;ret.f=ret.f%mod;
            ret.g+=((tmp.f*i)%mod*ten[dep-1]%mod+tmp.g)%mod;ret.g=ret.g%mod;
            ret.h+=(i*i%mod*ten[dep-1]%mod*ten[dep-1]%mod*tmp.f%mod+2*i*ten[dep-1]%mod*tmp.g%mod+tmp.h)%mod;ret.h=ret.h%mod;
        }
    }
    if(lim==0) f[dep][sta1][sta2]=ret;
    return ret;
}
long long calc(long long x){
    int cnt=0;
    while(x){c[++cnt]=x%10;x/=10;}
    return (dp(cnt,0,0,1).h)%mod;
}
int main(){
    ios::sync_with_stdio(0),cin.tie(0),cout.tie(0);
    ten[0]=1;
    for(int i=1;i<=17;i++) ten[i]=(ten[i-1]*10)%mod;
    cin>>n;memset(f,-1,sizeof f);
    for(int i=1;i<=n;i++){
        cin>>a>>b;
        cout<<(calc(b)-calc(a-1)+mod)%mod<<"\n";
    }
    return 0;
}

::::

Part 4

如何优化掉一个 B?例如 P3281 [SCOI2013] 数数。

我们先正常思考我们发现其实在我们枚举的时候有很大一部分情况的后记状态的答案是相同的,那么只需要用同一个 node tmp 来转移即可。 ::::info[code]

#include<bits/stdc++.h>
#define ll long long
using namespace std;
const ll N=1e5+5,mod=20130427;
struct node{
    ll ans,sum,cnt;
    //ans为答案
    //sum为以第i位开头的
}f[N][2];
ll B,n,m,a[N],s[N];
node dp(int dep,bool lim,bool qd0){
    if(dep==0) return {0,0,1};
    if(lim==0 && f[dep][qd0].cnt!=-1) return f[dep][qd0];
    ll up=lim?a[dep]:B-1;
    node res={0,0,0};
    for(int i=0;i<=up;i+=max(1ll,up)){
        node tmp=dp(dep-1,lim&&(i==up),qd0||(i!=0));
        if(!qd0&&!i){res.ans+=tmp.ans;res.ans%=mod;continue;}
        res.ans+=tmp.ans;res.ans%=mod;
        res.cnt+=tmp.cnt;res.cnt%=mod;
        res.sum+=tmp.sum+tmp.cnt*s[dep]%mod*i;res.sum%=mod;
    }
    node tmp=dp(dep-1,0,1);
    if(up>1){
        res.sum+=(up-1)*tmp.sum+tmp.cnt*s[dep]%mod*(up*(up-1)/2)%mod;res.sum%=mod;
        res.ans+=tmp.ans*(up-1);res.ans%=mod;
        res.cnt+=tmp.cnt*(up-1);res.cnt%=mod;
    }
    res.ans+=res.sum;res.ans%=mod;
    if(lim==0) f[dep][qd0]=res;
    return res;
}
ll calc(){
    memset(f,-1,sizeof f);
    return dp(n,1,0).ans;
}
signed main(){
    ios::sync_with_stdio(0),cout.tie(0),cin.tie(0);
    cin>>B;for(int i=1,pw=1;i<N;pw=pw*B%mod,i++) s[i]=(s[i-1]+pw)%mod;
    cin>>n;for(int i=n;i>=1;i--) cin>>a[i];
    a[1]--;
    for(int i=1;i<=n;i++){
        if(a[i]<0){
            a[i]+=B;
            a[i+1]--;
        }
    }
    ll ans=mod-calc();
    if(a[2]==-1) ans=0;
    cin>>n;for(int i=n;i>=1;i--) cin>>a[i];
    ans=ans+calc();
    cout<<ans%mod;
    return 0;
}

::::

Part 5

有限状态下的数位 dp。例如淘金对于一个位置 (i,j) 会变为 (f(i),f(j)),其中 f(i) 表示 i 数位上各位数的乘积,只能分解成 2^a \times 3^b \times 5^c \times 7^d 爆搜后发现有用的状态只有 14672 个,然后数位 dp 即可,最后用堆来维护 cnt(i) \times cnt(j) 的前 k 大值。 ::::info[code]

#include<bits/stdc++.h>
#define ll long long
using namespace std;
const ll mod=1e9+7;
ll n,k,c[15],f[15][40][25][15][15],prime[6]={0,2,3,5,7},cc,cnt,a[15000],b[15000];
//2 3 5 7
ll num[10][6]={
    {0,0,0,0,0},//0
    {0,0,0,0,0},//1
    {0,1,0,0,0},//2
    {0,0,1,0,0},//3
    {0,2,0,0,0},//4
    {0,0,0,1,0},//5
    {0,1,1,0,0},//6
    {0,0,0,0,1},//7
    {0,3,0,0,0},//8
    {0,0,2,0,0},//9
};
ll dp(int dep,int cnt2,int cnt3,int cnt5,int cnt7,bool lim,bool qd0){
    if(dep==0) return cnt2==0 && cnt3==0 && cnt5==0 && cnt7==0;
    if(lim==0 && qd0 && f[dep][cnt2][cnt3][cnt5][cnt7]!=-1) return f[dep][cnt2][cnt3][cnt5][cnt7];
    ll up=lim?c[dep]:9;
    ll res=0;
    for(int i=0;i<=up;i++){
        if(i==0 && qd0) continue;
        res+=dp(dep-1,cnt2-num[i][1],cnt3-num[i][2],cnt5-num[i][3],cnt7-num[i][4],lim&&(i==up),qd0||(i!=0));
    }
    if(lim==0 && qd0) f[dep][cnt2][cnt3][cnt5][cnt7]=res;
    return res;
}
void init(int pos,ll x){
    if(x>n) return ;
    a[++cc]=x;
    for(int i=pos;i<=4;i++){
        init(i,x*prime[i]);
    }
}
ll calc(ll x){
    ll cnt2=0,cnt3=0,cnt5=0,cnt7=0;
    while(x%2==0) cnt2++,x/=2;
    while(x%3==0) cnt3++,x/=3;
    while(x%5==0) cnt5++,x/=5;
    while(x%7==0) cnt7++,x/=7;
    return dp(cnt,cnt2,cnt3,cnt5,cnt7,1,0);
}
struct node{
    ll i,j;
    bool operator<(const node &o)const{
        return b[i]*b[j]<b[o.i]*b[o.j];
    }
};
int main(){
    ios::sync_with_stdio(0),cout.tie(0),cin.tie(0);
    cin>>n>>k;cnt=0;
    init(1,1);memset(f,-1,sizeof f);
    while(n){c[++cnt]=n%10;n/=10;}
    for(int i=1;i<=cc;i++) b[i]=calc(a[i]);b[1]--;
    sort(b+1,b+1+cc,greater<ll>());
    priority_queue<node> q;
    for(int i=1;i<=cc;i++) q.push({i,1});
    ll ans=0;
    while(k){
        node u=q.top();
        q.pop();
        ans=(ans+b[u.i]*b[u.j])%mod;
        k--;
        if(u.j!=cc) q.push({u.i,u.j+1});
    }
    cout<<ans;
    return 0;
}

::::

Part 6

不能直接数位 dp 的题目,需要算贡献或需要枚举一些值来帮助 dp。

例题 CF908G

直接计算不好计算,所以考虑每个数的贡献,我们发现可以将一个 $S(n)$ 看作一个全 $0$ 序列做后缀 $+1$,令 $sum(n)=\sum_{i=0}^{n-1} 10^i$ 那么一个数的 $S(n)=\sum_{i=1}^9 sum($ 数位中 $\ge i$ 的个数 $)$ 那么就可以做到 $O(n^2 B^2)$。 ::::info[code] ```cpp #include<bits/stdc++.h> #define ll long long using namespace std; const ll N=705,mod=1e9+7; ll n,f[N][N],c[N],pw[N],s[N],now; ll dp(int dep,int sum,int lim,int qd0){ if(dep==0) return s[sum]; if(!lim && qd0 && f[dep][sum]!=-1) return f[dep][sum]; ll up=lim?c[dep]:9; ll res=0; for(int i=0;i<=up;i++){ res+=dp(dep-1,sum+(i>=now),lim&&(i==up),qd0||(i!=0)); res%=mod; } if(!lim && qd0) f[dep][sum]=res; return res; } inline ll calc(string s){ ll cnt=0; for(int i=s.size()-1;i>=0;i--) c[++cnt]=s[i]-'0'; ll res=0; for(now=1;now<=9;now++){ memset(f,-1,sizeof f); res=(res+dp(cnt,0,1,0))%mod; } return res; } string x; signed main(){ ios::sync_with_stdio(0),cin.tie(0),cout.tie(0); pw[0]=1; for(int i=1;i<N;i++) s[i]=s[i-1]+pw[i-1],pw[i]=pw[i-1]*10%mod; cin>>x; cout<<calc(x); return 0; } ``` :::: # Part 7 多线 dp [例题](https://acm.hdu.edu.cn/showproblem.php?pid=5803),顾名思义就是并不是找到一个满足的数字而是找到数字对或三元组甚至四元组即以上,例如本题需要求的是 $0 \leq a \leq A,0 \leq b \leq B,0 \leq c \leq C,0 \leq d \leq D$ 中满足 $a+c > b+d \land a+d \ge b+c$ 的四元对有多少个。 考虑设计状态 $f_{dep,f1,f2,lim1,lim2,lim3,lim4}$ 直接记录 $f1=a+c-b-d > 0,f2 = a+d-b-c \ge 0$ 因为只有两个数所以只要大于二或小于等于二,那么无论如何后面的数都不会改变 $f1,f2$ 的正负情况,那么状态就为 $18 \times 5 \times 5 \times 2^4$ 但是每次枚举 $a,b,c,d$ 是需要 $10^4$ 直接转二进制优化即可。 ::::info[code] ```cpp #include<bits/stdc++.h> #define ll long long using namespace std; const ll N=1e6+5,mod=1e9+7; ll T,A,B,C,D; int a[61],b[61],c[61],d[61],f[61][6][6][2][2][2][2]; ll dp(int dep,int f1,int f2,int lim1,int lim2,int lim3,int lim4){ if(dep==0){ return f1>0&&f2>=0; } if(f[dep][f1+2][f2+2][lim1][lim2][lim3][lim4]!=-1) return f[dep][f1+2][f2+2][lim1][lim2][lim3][lim4]; ll up1=lim1?a[dep]:1,up2=lim2?b[dep]:1,up3=lim3?c[dep]:1,up4=lim4?d[dep]:1; ll res=0; for(int a=0;a<=up1;a++){ for(int b=0;b<=up2;b++){ for(int c=0;c<=up3;c++){ for(int d=0;d<=up4;d++){ ll zw1=max(-2,min(2,f1*2+a+c-b-d)),zw2=max(-2,min(2,f2*2+a+d-b-c)); if(zw1==-2 || zw2==-2) continue; res+=dp(dep-1,zw1,zw2,lim1&&(a==up1),lim2&&(b==up2),lim3&&(c==up3),lim4&&(d==up4)); if(res>=mod) res-=mod; } } } } f[dep][f1+2][f2+2][lim1][lim2][lim3][lim4]=res; return res; } ll calc(){ ll cnt=0; memset(f,-1,sizeof f); while(A || B || C || D){ a[++cnt]=A%2;A/=2; b[cnt]=B%2;B/=2; c[cnt]=C%2;C/=2; d[cnt]=D%2;D/=2; } return dp(cnt,0,0,1,1,1,1); } signed main(){ ios::sync_with_stdio(0),cin.tie(0),cout.tie(0); cin>>T; while(T--){ cin>>A>>B>>C>>D; cout<<calc()<<"\n"; } return 0; } ``` :::: # Part 8 [求回文数](https://www.luogu.com.cn/problem/B3883),先处理长度小于 $\vert n \vert$ 的回文数的个数,再使用数位 dp 顶着前 $\frac{|n|}{2}$ 位判断,枚举前面一半使得这个数一定为回文数即可。 # 练习题 ## [P5261 [JSOI2013] 数字理论](https://www.luogu.com.cn/problem/P5261) 非常有意思的一道题,为什么题解里面没有一个用递归做的,我们需要找到最小的 $n$ 位数的数位和为 $s1$ 且它乘 $d$ 以后的数位和为 $s2$。 我们考虑枚举它乘 $d$ 后的数,设计 dp 状态为 $f_{dep,d1,d2,lst}$ 表示还剩 $dep$ 位没填,该数的数位和还差 $d1$,$d$ 倍的该数的数位和还差 $d2$,模 $d$ 的余数为 $lst$。那么结束的状态就是 `!d1&&!d2&&!lst` 表示数位和恰好为 $s1,s2$ 且是 $d$ 的倍数,使用除法来推出它原本的数,详细见代码部分。 ::::info[code] ```cpp #include<bits/stdc++.h> #define ll long long using namespace std; const ll N=1e5+5; ll n,s1,s2,d,ans[N]; bitset<915> f[101][905][11]; void dfs(int dep,int d1,int d2,int lst){ if(d1<0||d2<0||d1>dep*9||d2>dep*9) return ; if(dep==0){ if(!d1&&!d2&&!lst){ for(int i=n;i>=1;i--) cout<<ans[i]; exit(0); } return ; } if(f[dep][d1][lst][d2]) return ; f[dep][d1][lst][d2]=1; for(int k=lst*10;k<lst*10+10;k++){ int i=k/d,j=k%d; ans[dep]=i; dfs(dep-1,d1-i,d2-k%10,j); } } inline ll calc(ll x){ ll res=0; while(x) res+=x%10,x/=10; return res; } signed main(){ ios::sync_with_stdio(0),cin.tie(0),cout.tie(0); cin>>n>>s1>>s2>>d; for(int i=1;i<=9;i++){ for(int j=0;j<d;j++){ ans[n]=i; dfs(n-1,s1-i,s2-calc(i*d+j),j); } } puts("-1"); return 0; } ``` :::: 开局先枚举原本的数的开头与模 $d$ 的余数,保证不包含前导零,时间复杂度 $O(10 \times KSP)$。 ## [P4067 [SDOI2016] 储能表](https://www.luogu.com.cn/problem/P4067) Part 3 + Part 7 的简单应用,需要求出: $$\sum_{i=0}^{n-1} \sum_{j=0}^{m-1} \max((i \oplus j)-k,0)$$ 先考虑如何算出: $$\sum_{i=0}^{n-1} \sum_{j=0}^{m-1} (i \oplus j)$$ 我们双线枚举 $n,m$,然后使用数位 dp 维护大于 $k$ 的个数与异或和,那么转移就是: ```cpp res.cnt=(res.cnt+tmp.cnt)%p; res.sum=(res.sum+tmp.sum+tmp.cnt*((x<<dep-1)%p))%p; ``` 这一位上的异或为 $x$ 再乘上 $cnt$ 就是到这一位的异或和,对于 $x$ 我们只需要再加一条限制,判断他是否到达下限 $k$ 即可。 ::::info[code] ```cpp #include<bits/stdc++.h> #define ll long long using namespace std; const ll N=65; struct node{ __int128 cnt,sum; }f[N][2][2][2]; ll T,n,m,t,p,a[N],b[N],c[N]; node dp(int dep,int limn,int limm,int limk){ if(dep==0) return {1,0}; if(f[dep][limn][limm][limk].cnt!=-1) return f[dep][limn][limm][limk]; ll upn=limn?a[dep]:1,upm=limm?b[dep]:1,dnk=limk?c[dep]:0; node res={0,0}; for(int i=0;i<=upn;i++){ for(int j=0;j<=upm;j++){ ll k=i^j; if(k>=dnk){ node tmp=dp(dep-1,limn&&(i==upn),limm&&(j==upm),limk&&(k==dnk)); res.cnt=(res.cnt+tmp.cnt)%p; res.sum=(res.sum+tmp.sum+tmp.cnt*((k<<dep-1)%p))%p; } } } f[dep][limn][limm][limk]=res; return res; } inline ll calc(ll n,ll m,ll x){ ll c1=0,c2=0,c3=0; memset(a,0,sizeof a); memset(b,0,sizeof b); memset(c,0,sizeof c); memset(f,-1,sizeof f); while(n){a[++c1]=n%2;n/=2;} while(m){b[++c2]=m%2;m/=2;} while(x){c[++c3]=x%2;x/=2;} node res=dp(max({c1,c2,c3}),1,1,1); // cout<<c1<<" "<<c2<<" "<<c3<<" "<<res.cnt<<" "<<res.sum<<"\n"; return (res.sum-res.cnt*t%p+p)%p; } signed main(){ ios::sync_with_stdio(0),cin.tie(0),cout.tie(0); cin>>T; while(T--){ cin>>n>>m>>t>>p;n--;m--; cout<<calc(n,m,t)<<"\n"; } return 0; } ``` :::: ## [P1149 [NOIP 2008 提高组] 火柴棒等式](https://www.luogu.com.cn/problem/P1149) 虽然这题的 $n$ 只有 $24$,所以我们决定大炮打蚊子,用 $O(n)$ 来解决这道题。 因为没有任何的限制,所以我们可以从低位开始枚举,因为这样好处理 $c$ 的进位,那么设置dp状态为 $f_{cnt,lima,limb,jw,fst}$ 因为要区分开前导零与零额外开一位 $fst$ 判断是否为最低位,然后这个状态表示的是已经使用了 $cnt$ 个木棍,$lima$ 若为 $1$ 表示数字 $a$ 已经锁定,$limb$ 同理,以及一维进位,~~我们就可以轻松爆标~~。 ::::info[code] ```cpp #include<bits/stdc++.h> #define ll long long using namespace std; const ll N=30,mod=1e9+7; ll f[N][2][2][2],n,num[]={6,2,5,5,4,5,6,3,7,6}; ll dp(ll cnt,int lima,int limb,int jw,int fst){ if(cnt>n) return 0; if(lima && limb) return cnt+(jw?2:0)==n; if(f[cnt][lima][limb][jw]!=-1) return f[cnt][lima][limb][jw]; ll upa=lima?0:9,upb=limb?0:9; ll res=0; for(int a=0;a<=upa;a++){ for(int b=0;b<=upb;b++){ ll c=a+b+jw; ll now=cnt+(!lima?num[a]:0)+(!limb?num[b]:0)+num[c%10]; res=(res+dp(now,lima,limb,c/10,0))%mod; if(!lima && (a || fst)) res=(res+dp(now,1,limb,c/10,0))%mod; if(!limb && (b || fst)) res=(res+dp(now,lima,1,c/10,0))%mod; if(!lima && (a || fst) && !limb && (b || fst)) res=(res+dp(now,1,1,c/10,0))%mod; } } f[cnt][lima][limb][jw]=res; return res; } signed main(){ ios::sync_with_stdio(0),cin.tie(0),cout.tie(0); cin>>n;n-=4; if(n<=0){return cout<<0,0;} memset(f,-1,sizeof f); cout<<dp(0,0,0,0,1); return 0; } ``` :::: ## [SP1433 KPSUM](https://www.luogu.com.cn/problem/SP1433) 这题就比较分类讨论了,长度为偶数的所有数第一位都是 `+`,若长度为奇数的其本身为奇则为 `+` 否则为 `-`。 ::::info[code] ```cpp #include<bits/stdc++.h> #define ll long long using namespace std; constexpr ll N=20; ll n,f[N][130][2][2][2][2],c[N]; ll dp(int dep,int sum,int lim,int qd0,int ws,int jo){ if(!qd0) ws=0; if(dep==0){ if(!ws) return 64-sum; else if(jo) return sum-64; else return 64-sum; } if(f[dep][sum][lim][qd0][ws][jo]!=-1) return f[dep][sum][lim][qd0][ws][jo]; ll res=0; ll up=lim?c[dep]:9; for(int i=0;i<=up;i++){ res+=dp(dep-1,sum+(ws%2==0?1:-1)*i,lim&&(i==up),qd0||(i!=0),ws^1,i&1); } f[dep][sum][lim][qd0][ws][jo]=res; return res; } inline ll calc(ll x){ memset(f,-1,sizeof f); ll cnt=0; while(x){c[++cnt]=x%10;x/=10;} return dp(cnt,64,1,0,0,0); } signed main(){ ios::sync_with_stdio(0),cin.tie(0),cout.tie(0); while(cin>>n){ if(!n) break; cout<<calc(n)<<"\n"; } return 0; } ``` :::: ## [2026hdu多校 1008FWT](https://acm.hdu.edu.cn/contest/problem?cid=1231&pid=1008) 可以列出一个递推版的数位 dp,类似于 **Sam 数**一样用矩阵快速幂来维护,且我们发现接下来的查询只会影响 $l,r$ 其中的一个的一位二进制上的数位,所以直接用线段树维护矩阵即可。打的表在代码下面,只要发现了性质这道题还是比较容易的。 ::::info[code] ```cpp #include<bits/stdc++.h> #define ll long long #define pll pair<ll,ll> using namespace std; const ll N=1e5+5,mod=998244353; ll T,n,m,q,cnt[5][5],len; struct matrix{ ll a[5][5],nn,mm; matrix(ll nn=0,ll mm=0):nn{nn},mm{mm}{memset(a,0,sizeof a);} matrix friend operator*(const matrix &x,const matrix &y){ matrix z(x.nn,y.mm); for(int i=0;i<x.nn;i++) for(int k=0;k<x.mm;k++) for(int j=0;j<y.mm;j++) z.a[i][j]=(z.a[i][j]+x.a[i][k]*y.a[k][j])%mod; return z; } matrix friend operator^(matrix &x,const ll&y){ matrix z(x.nn,x.mm),tmp=x; for(int i=0;i<x.nn;i++) z.a[i][i]=1; for(ll i=y;i;i>>=1){ if(i&1) z=z*tmp; tmp=tmp*tmp; } return z; } }; matrix tr[N<<2],leaf[N]; inline void pushup(int rt){ tr[rt]=tr[rt<<1]*tr[rt<<1|1]; } inline void build(int rt,int l,int r){ if(l==r){ tr[rt]=leaf[l]; return ; } int mid=l+r>>1; build(rt<<1,l,mid); build(rt<<1|1,mid+1,r); pushup(rt); } inline void modify(int rt,int l,int r,int x,matrix t){ if(l==r){ tr[rt]=t; return ; } int mid=l+r>>1; if(x<=mid) modify(rt<<1,l,mid,x,t); else modify(rt<<1|1,mid+1,r,x,t); pushup(rt); } inline ll query(){ matrix ans(1,4); ans.a[0][0]=1; ans=ans*tr[1]; ll res=0; for(int i=0;i<4;i++) res=(res+ans.a[0][i])%mod; return res; } vector<pll> xz; string l,r; signed main(){ ios::sync_with_stdio(0),cin.tie(0),cout.tie(0); cin>>T; while(T--){ cin>>n>>m>>q;xz.clear(); for(int i=0;i<2;i++) for(int j=0;j<2;j++) cnt[i][j]=0; for(int i=0;i<n;i++){ if(i) xz.push_back({0,i}); if(i!=n-1) xz.push_back({i,n-1}); } for(int i=1;i<=m;i++){ ll x,y;cin>>x>>y; x--;y--; xz.push_back({x,y}); } for(int s=0;s<(1<<n);s++){ bool falg=1; for(auto i:xz){ int u=i.first,v=i.second; u=(s>>u)&1,v=(s>>v)&1; if(u<v){falg=0;break;} } if(falg){ cnt[(s>>(n-1))&1][s&1]++; } } cin>>l>>r; ll t1=l.size(),t2=r.size(); len=max(l.size(),r.size()); string t=""; if(len==l.size()){for(int i=1;i<=len-(ll)r.size();i++) t+="0";r=t+r;} else{ for(int i=1;i<=len-(ll)l.size();i++) t+="0"; l=t+l; } for(int i=0;i<len;i++){ matrix M(4,4); int lb=l[i]-'0',rb=r[i]-'0'; for(int a=0;a<2;a++){ for(int b=0;b<2;b++){ if(cnt[a][b]==0) continue; for(int st=0;st<4;st++){ int liml=(st>>1)&1; int limr=st&1;//0未达上限 if(liml==0 && a<lb) continue; if(limr==0 && b>rb) continue; int nliml,nlimr; if(liml==1) nliml=1; else if(a>lb) nliml=1; else nliml=0; if(limr==1) nlimr=1; else if(b<rb) nlimr=1; else nlimr=0; int nst=(nliml<<1)|nlimr; M.a[st][nst]=(M.a[st][nst]+cnt[a][b])%mod; } } } leaf[i]=M; } build(1,0,len-1); cout<<query()<<"\n"; while(q--){ int op,pos; cin>>op>>pos;pos--; if(op==0) pos+=len-t1; else pos+=len-t2; string &s=(op==0?l:r); (s[pos]=='0'?s[pos]='1':s[pos]='0'); matrix M(4,4); int lb=l[pos]-'0',rb=r[pos]-'0'; for(int a=0;a<2;a++){ for(int b=0;b<2;b++){ if(cnt[a][b]==0) continue; for(int st=0;st<4;st++){ int liml=(st>>1)&1; int limr=st&1;//0未达上限 if(liml==0 && a<lb) continue; if(limr==0 && b>rb) continue; int nliml,nlimr; if(liml==1) nliml=1; else if(a>lb) nliml=1; else nliml=0; if(limr==1) nlimr=1; else if(b<rb) nlimr=1; else nlimr=0; int nst=(nliml<<1)|nlimr; M.a[st][nst]=(M.a[st][nst]+cnt[a][b])%mod; } } } modify(1,0,len-1,pos,M); cout<<query()<<"\n"; } } return 0; } /* l <= xl <= xr <= r dp[dep][liml][limr]=dp[dep-1][] x1&xi=xi && xi&xn=xn a b c a&b=b b&c=c a b c 1/0 0 0 1 1 1/0 a>=b>=c x1>=xi>=xn a&b==b a b 1/0 0 1 1 a>=b xa>=xb */ ``` :::: ## [P5674 「SWTR-2」Magical Gates](https://www.luogu.com.cn/problem/P5674) 如果输入的 $p$ 并不是质数,不好求逆元,那么将可以使用 Part 7 的变形,上下界 dp 与上一题 FWT 相同,需要求出 $\sum_l^r d_i \mod p$ 与 $\prod_l^r d_i \mod p$。 那么直接设置 dp 状态为 $f_{dep,sum,lim1,lim2}$ 分别表示已经出现过的 $1$ 的个数,是否达到上界,是否达到下界,然后直接数位 dp 即可,但是最麻烦的部分在于高精度。 ::::info[逆元code] ```cpp #include<bits/stdc++.h> #define ll long long using namespace std; constexpr ll M=1000/8+5,inf=1e18; const int base=1e8; int aux[M<<3]; struct bignum{ int s[M],l; bool neg; inline void clr(){l=0;memset(s,0,sizeof(s));neg=false;}bignum(){clr();} inline void write(){if(neg&&!(l==1&&s[1]==0))putchar('-');printf("%d",s[l]);for(int i=l-1;i;i--)printf("%08d",s[i]);} inline void read(){char c=getchar();bool sign=false;while(c!='-'&&(c<'0'||c>'9'))c=getchar();if(c=='-'){sign=true;c=getchar();}int i,x=0,k=1,L=0,fl,o;for(;c>='0'&&c<='9';c=getchar()){if(!(L-1)&&!aux[L])L--;aux[++L]=c-'0';}clr();neg=sign;l=L/8+((o=L%8)>0);for(i=1;i<=o;i++)x=x*10+aux[i];if(o)s[l]=x;for(fl=!o?l+1:l,i=o+1,x=0;i<=L;i++,k++){x=x*10+aux[i];if(!(k^8))s[--fl]=x,x=k=0;}if(!l)l=1;if(l==1&&s[1]==0)neg=false;} inline ll toint()const{ll x=0;for(int i=l;i;i--)x=x*base+s[i];return neg?-x:x;} inline bignum& operator=(int b){clr();if(b<0){neg=true;b=-b;}do{s[++l]=b%base;b/=base;}while(b>0);return *this;} inline bignum& operator=(ll b){clr();if(b<0){neg=true;b=-b;}do{s[++l]=b%base;b/=base;}while(b>0);return *this;} inline bool abs_less(const bignum& b)const{if(l^b.l)return l<b.l;for(int i=l;i;i--)if(s[i]^b.s[i])return s[i]<b.s[i];return false;} inline bignum abs_add(const bignum& b)const{bignum c;c.clr();ll x=0;int k=max(l,b.l);c.l=k;for(int i=1;i<=k;i++){x=x+(i<=l?s[i]:0)+(i<=b.l?b.s[i]:0);c.s[i]=x%base;x/=base;}if(x)c.s[++c.l]=x;return c;} inline bignum abs_sub(const bignum& b)const{bignum c,d=*this;c.clr();ll x=0;for(int i=1;i<=l;i++){if((x=d.s[i])<b.s[i]){d.s[i+1]--;x+=base;}c.s[i]=x-b.s[i];}c.l=l;for(;!c.s[c.l]&&c.l>1;c.l--);return c;} inline bignum operator-()const{bignum c=*this;if(!(l==1&&s[1]==0))c.neg=!neg;return c;} inline bignum operator+(const bignum& b)const{if(neg==b.neg){bignum c=abs_add(b);c.neg=neg;if(c.l==1&&c.s[1]==0)c.neg=false;return c;}else{if(abs_less(b)){bignum c=b.abs_sub(*this);c.neg=b.neg;if(c.l==1&&c.s[1]==0)c.neg=false;return c;}else{bignum c=abs_sub(b);c.neg=neg;if(c.l==1&&c.s[1]==0)c.neg=false;return c;}}} inline bignum operator-(const bignum& b)const{bignum neg_b=b;neg_b.neg=!b.neg;return *this+neg_b;} inline bignum operator*(const bignum& b)const{if(b.l<2)return *this*b.toint();bignum c;ll x;int i,j,k;c.clr();for(i=1;i<=l;i++){x=0;for(j=1;j<=b.l;j++){x=x+1LL*s[i]*b.s[j]+c.s[k=i+j-1];c.s[k]=x%base;x/=base;}if(x)c.s[i+b.l]=x;}c.l=l+b.l;for(;!c.s[c.l]&&c.l>1;c.l--);c.neg=neg^b.neg;if(c.l==1&&c.s[1]==0)c.neg=false;return c;} inline bignum operator*(const ll& b)const{ll bb=b<0?-b:b;bignum c;c.clr();ll x=0;for(int i=1;i<=l;i++){x=x+1LL*s[i]*bb;c.s[i]=x%base;x/=base;}c.l=l;while(x){c.s[++c.l]=x%base;x/=base;}c.neg=neg^(b<0);if(c.l==1&&c.s[1]==0)c.neg=false;return c;} inline bignum operator*(const int& b)const{return *this*(ll)b;} inline bignum operator/(const bignum& b)const{if(b.l<2)return *this/b.toint();bignum c,d;int i,j,le,r,mid,k;c.clr();d.clr();bignum abs_this=*this,abs_b=b;abs_this.neg=abs_b.neg=false;for(i=l;i;i--){for(j=++d.l;j>1;j--)d.s[j]=d.s[j-1];d.s[1]=abs_this.s[i];if(d<abs_b)continue;le=k=0;r=base-1;while(le<=r){mid=(le+r)>>1;bignum tmp=abs_b*mid;(tmp<=d)?(le=mid+1,k=mid):(r=mid-1);}c.s[i]=k;d=d-abs_b*k;}c.l=l;for(;!c.s[c.l]&&c.l>1;c.l--);c.neg=neg^b.neg;if(c.l==1&&c.s[1]==0)c.neg=false;return c;} inline bignum operator%(const bignum& b)const{if(b.l<2)return *this%b.toint();bignum c;int i,j,le,r,mid,k;c.clr();bignum abs_this=*this,abs_b=b;abs_this.neg=abs_b.neg=false;for(i=l;i;i--){for(j=++c.l;j>1;j--)c.s[j]=c.s[j-1];c.s[1]=abs_this.s[i];if(c<abs_b)continue;le=k=0;r=base-1;while(le<=r){mid=(le+r)>>1;bignum tmp=abs_b*mid;(tmp<=c)?(le=mid+1,k=mid):(r=mid-1);}c=c-abs_b*k;}for(;!c.s[c.l]&&c.l>1;c.l--);c.neg=neg;if(c.l==1&&c.s[1]==0)c.neg=false;return c;} inline bignum operator+(const int& b)const{bignum tmp;tmp=b;return *this+tmp;} inline bignum operator+(const ll& b)const{bignum tmp;tmp=b;return *this+tmp;} inline bignum operator-(const int& b)const{bignum tmp;tmp=b;return *this-tmp;} inline bignum operator-(const ll& b)const{bignum tmp;tmp=b;return *this-tmp;} inline bignum operator/(const int& b)const{bignum c;ll x=0;c.clr();for(int i=l;i;i--){c.s[i]=(x*base+s[i])/b;x=(x*base+s[i])%b;}for(c.l=l;!c.s[c.l]&&c.l>1;c.l--);return c;} inline bignum operator/(const ll & b)const{bignum c;ll x=0;c.clr();for(int i=l;i;i--){c.s[i]=(x*base+s[i])/b;x=(x*base+s[i])%b;}for(c.l=l;!c.s[c.l]&&c.l>1;c.l--);return c;} inline bignum operator%(const int& b)const{bignum c;ll x=0;c.clr();for(int i=l;i;i--)x=(x*base+s[i])%b;return c=x;} inline bignum operator%(const ll & b)const{bignum c;ll x=0;c.clr();for(int i=l;i;i--)x=(x*base+s[i])%b;return c=x;} inline bignum operator+=(const bignum& b){return *this=*this+b;} inline bignum operator+=(int b){bignum tmp;tmp=b;return *this+=tmp;} inline bignum operator+=(ll b){bignum tmp;tmp=b;return *this+=tmp;} inline bignum operator-=(const bignum& b){return *this=*this-b;} inline bignum operator-=(int b){bignum tmp;tmp=b;return *this-=tmp;} inline bignum operator-=(ll b){bignum tmp;tmp=b;return *this-=tmp;} inline bignum operator*=(const bignum& b){return *this=*this*b;} inline bignum operator*=(int b){return *this=*this*(ll)b;} inline bignum operator*=(ll b){return *this=*this*b;} inline bignum operator/=(const bignum& b){return *this=*this/b;} inline bignum operator/=(int b){bignum tmp;tmp=b;return *this/=tmp;} inline bignum operator/=(ll b){bignum tmp;tmp=b;return *this/=tmp;} inline bignum operator%=(const bignum& b){return *this=*this%b;} inline bignum operator%=(int b){bignum tmp;tmp=b;return *this%=tmp;} inline bignum operator%=(ll b){bignum tmp;tmp=b;return *this%=tmp;} inline bool operator<(const bignum& b)const{if(neg!=b.neg)return neg>b.neg;if(neg==false)return abs_less(b);else return b.abs_less(*this);} inline bool operator<=(const bignum& b)const{return !(*this>b);} inline bool operator>(const bignum& b)const{return b<*this;} inline bool operator>=(const bignum& b)const{return !(*this<b);} inline bool operator==(const bignum& b)const{if(neg!=b.neg)return false;if(l!=b.l)return false;for(int i=l;i;i--)if(s[i]!=b.s[i])return false;return true;} inline bool operator!=(const bignum& b)const{return !(*this==b);} inline bool operator<(int b)const{bignum tmp;tmp=b;return *this<tmp;} inline bool operator<=(int b)const{bignum tmp;tmp=b;return *this<=tmp;} inline bool operator>(int b)const{bignum tmp;tmp=b;return *this>tmp;} inline bool operator>=(int b)const{bignum tmp;tmp=b;return *this>=tmp;} inline bool operator==(int b)const{bignum tmp;tmp=b;return *this==tmp;} inline bool operator!=(int b)const{bignum tmp;tmp=b;return *this!=tmp;} inline bool operator<(ll b)const{bignum tmp;tmp=b;return *this<tmp;} inline bool operator<=(ll b)const{bignum tmp;tmp=b;return *this<=tmp;} inline bool operator>(ll b)const{bignum tmp;tmp=b;return *this>tmp;} inline bool operator>=(ll b)const{bignum tmp;tmp=b;return *this>=tmp;} inline bool operator==(ll b)const{bignum tmp;tmp=b;return *this==tmp;} inline bool operator!=(ll b)const{bignum tmp;tmp=b;return *this!=tmp;} }l,r; constexpr ll N=4005; ll T,mod,wk,f1[N][N],f2[N][N],c1[N],c2[N]; ll dp1(int dep,int sum,int lim,int op){ if(dep==0) return sum; if(!lim && f1[dep][sum]!=-1) return f1[dep][sum]; ll up=lim?op?c1[dep]:c2[dep]:1,res=0; for(int i=0;i<=up;i++){ (res+=dp1(dep-1,sum+i,lim&&(i==up),op))%=mod; } if(!lim) f1[dep][sum]=res; return res; } ll dp2(int dep,int sum,int lim,int op){ if(dep==0) return max(sum,1); if(!lim && f2[dep][sum]!=-1) return f2[dep][sum]; ll up=lim?op?c1[dep]:c2[dep]:1,res=1; for(int i=0;i<=up;i++){ (res*=dp2(dep-1,sum+i,lim&&(i==up),op))%=mod; } if(!lim) f2[dep][sum]=res; return res; } ll ksm(ll a,ll b){ll res=1;while(b){if(b&1) res=res*a%mod;a=a*a%mod;b>>=1;}return res;} signed main(){ scanf("%lld %lld %lld",&T,&mod,&wk); memset(f1,-1,sizeof f1); memset(f2,-1,sizeof f2); while(T--){ l.read();r.read(); l-=1; ll cnt1=0,cnt2=0; while(r>0){ c1[++cnt1]=r.s[1]%2; r/=2; } while(l>0){ c2[++cnt2]=l.s[1]%2; l/=2; } printf("%lld %lld\n",(dp1(cnt1,0,1,1)-dp1(cnt2,0,1,0)+mod)%mod,dp2(cnt1,0,1,1)*ksm(dp2(cnt2,0,1,0),mod-2)%mod); } return 0; } ``` :::: ::::info[上下界code] ```cpp #include<bits/stdc++.h> #define ll long long using namespace std; constexpr ll M=1000/8+5,inf=1e18; const int base=1e8; int aux[M<<3]; struct bignum{ int s[M],l; bool neg; inline void clr(){l=0;memset(s,0,sizeof(s));neg=false;}bignum(){clr();} inline void write(){if(neg&&!(l==1&&s[1]==0))putchar('-');printf("%d",s[l]);for(int i=l-1;i;i--)printf("%08d",s[i]);} inline void read(){char c=getchar();bool sign=false;while(c!='-'&&(c<'0'||c>'9'))c=getchar();if(c=='-'){sign=true;c=getchar();}int i,x=0,k=1,L=0,fl,o;for(;c>='0'&&c<='9';c=getchar()){if(!(L-1)&&!aux[L])L--;aux[++L]=c-'0';}clr();neg=sign;l=L/8+((o=L%8)>0);for(i=1;i<=o;i++)x=x*10+aux[i];if(o)s[l]=x;for(fl=!o?l+1:l,i=o+1,x=0;i<=L;i++,k++){x=x*10+aux[i];if(!(k^8))s[--fl]=x,x=k=0;}if(!l)l=1;if(l==1&&s[1]==0)neg=false;} inline ll toint()const{ll x=0;for(int i=l;i;i--)x=x*base+s[i];return neg?-x:x;} inline bignum& operator=(int b){clr();if(b<0){neg=true;b=-b;}do{s[++l]=b%base;b/=base;}while(b>0);return *this;} inline bignum& operator=(ll b){clr();if(b<0){neg=true;b=-b;}do{s[++l]=b%base;b/=base;}while(b>0);return *this;} inline bool abs_less(const bignum& b)const{if(l^b.l)return l<b.l;for(int i=l;i;i--)if(s[i]^b.s[i])return s[i]<b.s[i];return false;} inline bignum abs_add(const bignum& b)const{bignum c;c.clr();ll x=0;int k=max(l,b.l);c.l=k;for(int i=1;i<=k;i++){x=x+(i<=l?s[i]:0)+(i<=b.l?b.s[i]:0);c.s[i]=x%base;x/=base;}if(x)c.s[++c.l]=x;return c;} inline bignum abs_sub(const bignum& b)const{bignum c,d=*this;c.clr();ll x=0;for(int i=1;i<=l;i++){if((x=d.s[i])<b.s[i]){d.s[i+1]--;x+=base;}c.s[i]=x-b.s[i];}c.l=l;for(;!c.s[c.l]&&c.l>1;c.l--);return c;} inline bignum operator-()const{bignum c=*this;if(!(l==1&&s[1]==0))c.neg=!neg;return c;} inline bignum operator+(const bignum& b)const{if(neg==b.neg){bignum c=abs_add(b);c.neg=neg;if(c.l==1&&c.s[1]==0)c.neg=false;return c;}else{if(abs_less(b)){bignum c=b.abs_sub(*this);c.neg=b.neg;if(c.l==1&&c.s[1]==0)c.neg=false;return c;}else{bignum c=abs_sub(b);c.neg=neg;if(c.l==1&&c.s[1]==0)c.neg=false;return c;}}} inline bignum operator-(const bignum& b)const{bignum neg_b=b;neg_b.neg=!b.neg;return *this+neg_b;} inline bignum operator*(const bignum& b)const{if(b.l<2)return *this*b.toint();bignum c;ll x;int i,j,k;c.clr();for(i=1;i<=l;i++){x=0;for(j=1;j<=b.l;j++){x=x+1LL*s[i]*b.s[j]+c.s[k=i+j-1];c.s[k]=x%base;x/=base;}if(x)c.s[i+b.l]=x;}c.l=l+b.l;for(;!c.s[c.l]&&c.l>1;c.l--);c.neg=neg^b.neg;if(c.l==1&&c.s[1]==0)c.neg=false;return c;} inline bignum operator*(const ll& b)const{ll bb=b<0?-b:b;bignum c;c.clr();ll x=0;for(int i=1;i<=l;i++){x=x+1LL*s[i]*bb;c.s[i]=x%base;x/=base;}c.l=l;while(x){c.s[++c.l]=x%base;x/=base;}c.neg=neg^(b<0);if(c.l==1&&c.s[1]==0)c.neg=false;return c;} inline bignum operator*(const int& b)const{return *this*(ll)b;} inline bignum operator/(const bignum& b)const{if(b.l<2)return *this/b.toint();bignum c,d;int i,j,le,r,mid,k;c.clr();d.clr();bignum abs_this=*this,abs_b=b;abs_this.neg=abs_b.neg=false;for(i=l;i;i--){for(j=++d.l;j>1;j--)d.s[j]=d.s[j-1];d.s[1]=abs_this.s[i];if(d<abs_b)continue;le=k=0;r=base-1;while(le<=r){mid=(le+r)>>1;bignum tmp=abs_b*mid;(tmp<=d)?(le=mid+1,k=mid):(r=mid-1);}c.s[i]=k;d=d-abs_b*k;}c.l=l;for(;!c.s[c.l]&&c.l>1;c.l--);c.neg=neg^b.neg;if(c.l==1&&c.s[1]==0)c.neg=false;return c;} inline bignum operator%(const bignum& b)const{if(b.l<2)return *this%b.toint();bignum c;int i,j,le,r,mid,k;c.clr();bignum abs_this=*this,abs_b=b;abs_this.neg=abs_b.neg=false;for(i=l;i;i--){for(j=++c.l;j>1;j--)c.s[j]=c.s[j-1];c.s[1]=abs_this.s[i];if(c<abs_b)continue;le=k=0;r=base-1;while(le<=r){mid=(le+r)>>1;bignum tmp=abs_b*mid;(tmp<=c)?(le=mid+1,k=mid):(r=mid-1);}c=c-abs_b*k;}for(;!c.s[c.l]&&c.l>1;c.l--);c.neg=neg;if(c.l==1&&c.s[1]==0)c.neg=false;return c;} inline bignum operator+(const int& b)const{bignum tmp;tmp=b;return *this+tmp;} inline bignum operator+(const ll& b)const{bignum tmp;tmp=b;return *this+tmp;} inline bignum operator-(const int& b)const{bignum tmp;tmp=b;return *this-tmp;} inline bignum operator-(const ll& b)const{bignum tmp;tmp=b;return *this-tmp;} inline bignum operator/(const int& b)const{bignum c;ll x=0;c.clr();for(int i=l;i;i--){c.s[i]=(x*base+s[i])/b;x=(x*base+s[i])%b;}for(c.l=l;!c.s[c.l]&&c.l>1;c.l--);return c;} inline bignum operator/(const ll & b)const{bignum c;ll x=0;c.clr();for(int i=l;i;i--){c.s[i]=(x*base+s[i])/b;x=(x*base+s[i])%b;}for(c.l=l;!c.s[c.l]&&c.l>1;c.l--);return c;} inline bignum operator%(const int& b)const{bignum c;ll x=0;c.clr();for(int i=l;i;i--)x=(x*base+s[i])%b;return c=x;} inline bignum operator%(const ll & b)const{bignum c;ll x=0;c.clr();for(int i=l;i;i--)x=(x*base+s[i])%b;return c=x;} inline bignum operator+=(const bignum& b){return *this=*this+b;} inline bignum operator+=(int b){bignum tmp;tmp=b;return *this+=tmp;} inline bignum operator+=(ll b){bignum tmp;tmp=b;return *this+=tmp;} inline bignum operator-=(const bignum& b){return *this=*this-b;} inline bignum operator-=(int b){bignum tmp;tmp=b;return *this-=tmp;} inline bignum operator-=(ll b){bignum tmp;tmp=b;return *this-=tmp;} inline bignum operator*=(const bignum& b){return *this=*this*b;} inline bignum operator*=(int b){return *this=*this*(ll)b;} inline bignum operator*=(ll b){return *this=*this*b;} inline bignum operator/=(const bignum& b){return *this=*this/b;} inline bignum operator/=(int b){bignum tmp;tmp=b;return *this/=tmp;} inline bignum operator/=(ll b){bignum tmp;tmp=b;return *this/=tmp;} inline bignum operator%=(const bignum& b){return *this=*this%b;} inline bignum operator%=(int b){bignum tmp;tmp=b;return *this%=tmp;} inline bignum operator%=(ll b){bignum tmp;tmp=b;return *this%=tmp;} inline bool operator<(const bignum& b)const{if(neg!=b.neg)return neg>b.neg;if(neg==false)return abs_less(b);else return b.abs_less(*this);} inline bool operator<=(const bignum& b)const{return !(*this>b);} inline bool operator>(const bignum& b)const{return b<*this;} inline bool operator>=(const bignum& b)const{return !(*this<b);} inline bool operator==(const bignum& b)const{if(neg!=b.neg)return false;if(l!=b.l)return false;for(int i=l;i;i--)if(s[i]!=b.s[i])return false;return true;} inline bool operator!=(const bignum& b)const{return !(*this==b);} inline bool operator<(int b)const{bignum tmp;tmp=b;return *this<tmp;} inline bool operator<=(int b)const{bignum tmp;tmp=b;return *this<=tmp;} inline bool operator>(int b)const{bignum tmp;tmp=b;return *this>tmp;} inline bool operator>=(int b)const{bignum tmp;tmp=b;return *this>=tmp;} inline bool operator==(int b)const{bignum tmp;tmp=b;return *this==tmp;} inline bool operator!=(int b)const{bignum tmp;tmp=b;return *this!=tmp;} inline bool operator<(ll b)const{bignum tmp;tmp=b;return *this<tmp;} inline bool operator<=(ll b)const{bignum tmp;tmp=b;return *this<=tmp;} inline bool operator>(ll b)const{bignum tmp;tmp=b;return *this>tmp;} inline bool operator>=(ll b)const{bignum tmp;tmp=b;return *this>=tmp;} inline bool operator==(ll b)const{bignum tmp;tmp=b;return *this==tmp;} inline bool operator!=(ll b)const{bignum tmp;tmp=b;return *this!=tmp;} }l,r; constexpr ll N=4005; ll T,mod,wk,f1[N][N],f2[N][N],c1[N],c2[N]; ll dp1(int dep,int sum,int lim1,int lim2){ if(dep==0) return sum; if(!lim1 && !lim2 && f1[dep][sum]!=-1) return f1[dep][sum]; ll up=lim1?c1[dep]:1,dn=lim2?c2[dep]:0,res=0; for(int i=dn;i<=up;i++){ (res+=dp1(dep-1,sum+i,lim1&&(i==up),lim2&&(i==dn)))%=mod; } if(!lim1 && !lim2) f1[dep][sum]=res; return res; } ll dp2(int dep,int sum,int lim1,int lim2){ if(dep==0) return sum; if(!lim1 && !lim2 && f2[dep][sum]!=-1) return f2[dep][sum]; ll up=lim1?c1[dep]:1,dn=lim2?c2[dep]:0,res=1; for(int i=dn;i<=up;i++){ (res*=dp2(dep-1,sum+i,lim1&&(i==up),lim2&&(i==dn)))%=mod; } if(!lim1 && !lim2) f2[dep][sum]=res; return res; } signed main(){ scanf("%lld %lld %lld",&T,&mod,&wk); memset(f1,-1,sizeof f1); memset(f2,-1,sizeof f2); while(T--){ l.read();r.read(); ll cnt1=0,cnt2=0; while(r>0){ c1[++cnt1]=r.s[1]%2; r/=2; } while(l>0){ c2[++cnt2]=l.s[1]%2; l/=2; } ll cnt=max(cnt1,cnt2); for(int i=cnt1+1;i<=cnt;i++) c1[i]=0; for(int i=cnt2+1;i<=cnt;i++) c2[i]=0; printf("%lld %lld\n",dp1(cnt,0,1,1),dp2(cnt,0,1,1)); } return 0; } ``` :::: ## [P8688 [蓝桥杯 2019 省 A] 组合数问题](https://www.luogu.com.cn/problem/P8688) 模板 lucas + 数位 dp 练习题,用到了 Part 4 + Part 7,首先我们先考虑如何暴力,对于组合数模一个质数想到 lucas 定理得到。 $$\binom{a}{b} \mod k = \binom{a/k}{b/k} \binom{a \mod k}{b \mod k} \mod k$$ 那么递归下去就得到了 $$\binom{a}{b} \mod k = \binom{a_1}{b_1} \cdots \binom{a_n}{b_n} \mod k$$ 其中 $\begin{cases} a_1 + k a_2 + k^2 a_3 + \cdots + k^{n-1} a_n = a\\b_1 + k b_2 + k^2 b_3 + \cdots + k^{n-1} b_n = b \end{cases}

且因为 a_i,b_i < k 所以 \binom{a_i}{b_i} \equiv 0 \mod k 当且仅当 a_i < b_i

稍微容斥一下就是总方案数减去不为 0 的方案数也就是所有的 a_i \ge b_i 的方案数,那么就变为了 k 进制 dp,1 \leq a \leq n,0 \leq b \leq m 有多少对这样的数对就可以写出一个暴力的版本,时间复杂度 O(\log_k n \times k^2) 代码如下,可以通过弱化版 P6669。 ::::info[暴力]

#include<bits/stdc++.h>
using namespace std;
#define ll long long
constexpr ll N=65,mod=1e9+7;
ll T,k,n,m,f[N][2][2],c1[N],c2[N];
ll dp(int dep,int lim1,int lim2){
    if(dep==0) return 1;
    if(f[dep][lim1][lim2]!=-1) return f[dep][lim1][lim2];
    int up1=lim1?c1[dep]:k-1,up2=lim2?c2[dep]:k-1;
    ll res=0;
    for(int a=0;a<=up1;a++){
        for(int b=0;b<=min(up2,a);b++){
            res+=dp(dep-1,lim1&&(a==up1),lim2&&(b==up2));
            res%=mod;
        }
    }
    f[dep][lim1][lim2]=res;
    return res;
}
signed main(){
    ios::sync_with_stdio(0),cout.tie(0),cin.tie(0);
    cin>>T>>k;
    while(T--){
        cin>>n>>m;
        m=min(n,m);
        ll ans=(n%mod*((m+1)%mod)%mod-(m%mod)*((m-1)%mod)%mod*500000004ll%mod+mod)%mod;
        ll cnt1=0,cnt2=0;
        while(n){
            c1[++cnt1]=n%k;
            n/=k;
        }
        while(m){
            c2[++cnt2]=m%k;
            m/=k;
        }
        memset(f,-1,sizeof f);
        for(int i=cnt2+1;i<=cnt1;i++) c2[i]=0;
        cout<<(ans-dp(cnt1,1,1)+mod+1)%mod<<"\n";
    }
    return 0;
}

:::: 写起来细节还是比较多的,需要加一是因为会把 (0,0) 多减了,再考虑如何使用 Part 4 来优化,那么我们就记录一个状态他的后继状态如 dp(dep-1,1,1) 出现了几次,然后分类讨论即可,该取模时就取模。 ::::info[code]

#include<bits/stdc++.h>
using namespace std;
#define ll long long
constexpr ll N=65,mod=1e9+7,inv2=500000004;
ll T,k,n,m,f[N][2][2],c1[N],c2[N];
ll dp(int dep,int lim1,int lim2){
    if(dep==0) return 1;
    if(f[dep][lim1][lim2]!=-1) return f[dep][lim1][lim2];
    ll up1=lim1?c1[dep]:k-1,up2=lim2?c2[dep]:k-1;
    // ll res=0;
    // for(int a=0;a<=up1;a++){
    //  for(int b=0;b<=min(up2,a);b++){
    //      res+=dp(dep-1,lim1&&(a==up1),lim2&&(b==up2));
    //      res%=mod;
    //  }
    // }
    ll t11=dp(dep-1,1,1),t10=dp(dep-1,1,0),t01=dp(dep-1,0,1),t00=dp(dep-1,0,0);
    ll tot=(up2>=up1?(2+up1)*(1+up1)%mod*inv2%mod:((2+up2)*(1+up2)%mod*inv2%mod+(up2+1)*(up1-up2)%mod)%mod);
    ll c11=0,c10=0,c01=0,c00=0;
    if(lim1 && lim2){
        c11=up1>=up2?1:0;
        c10=min(up1+1,up2);
        c01=up1>=up2?up1-up2:0;
        c00=tot-c11-c10-c01;
    }else if(lim1){
        c10=min(up1,up2)+1;
        c00=tot-c10;
    }else if(lim2){
        c01=up1>=up2?up1-up2+1:0;
        c00=tot-c01;
    }else c00=tot;
    c11=(c11%mod+mod)%mod;c10=(c10%mod+mod)%mod;
    c01=(c01%mod+mod)%mod;c00=(c00%mod+mod)%mod;
    f[dep][lim1][lim2]=(t11*c11+t10*c10+t01*c01+t00*c00)%mod;
    return f[dep][lim1][lim2];
}
signed main(){
    ios::sync_with_stdio(0),cout.tie(0),cin.tie(0);
    cin>>T>>k;
    while(T--){
        cin>>n>>m;
        m=min(n,m);
        ll ans=(n%mod*((m+1)%mod)%mod-(m%mod)*((m-1)%mod)%mod*inv2%mod+mod)%mod;
        ll cnt1=0,cnt2=0;
        while(n){
            c1[++cnt1]=n%k;
            n/=k;
        }
        while(m){
            c2[++cnt2]=m%k;
            m/=k;
        }
        for(int i=0;i<=cnt1;i++) f[i][1][1]=f[i][0][1]=f[i][1][0]=f[i][0][0]=-1;
        for(int i=cnt2+1;i<=cnt1;i++) c2[i]=0;
        cout<<(ans-dp(cnt1,1,1)+mod+1)%mod<<"\n";
    }
    return 0;
}

::::

P13125 [GCJ 2019 Finals] Won't sum? Must now

根据附件给的题解模拟的,但是时间复杂度也是错的,剪枝草过去的(也有可能是我时间复杂度分析错了),需要写出一种能高效判断是否存在 S=A+BAB 均为回文数,那么不妨设 A \ge B,那么 A 的位数只能是 |S||S|-1,再枚举 B 的位数,时间复杂度为 O(n)

然后最麻烦的是考虑进位,我们搜索时需要记录 (l,lc,rc) 分别表示正在处理的低位 l,(下表从 0 开始),r=la-1-l 表示高位的下表,lc 表示从 ll+1 位的进位,rc 表示 rr+1 位的进位。

结束条件:如果 l>r 那么判断 lc 是否等于 rc 即可也就是最后一次的进位是否相同,如果 l==r 那么 a_l=a_r 枚举 da 表示这一位的数值为多少根据等式:

da+db+lc=s_l+10*rc

对于一般的,枚举 nlcnrc 表示下一个 lcrc 只需要满足,低位:(nlcl 位向 l+1 位)

a_l+b_l+lc=s_l+10*nlc

高位:(nrcr-1 位向 r 位)

a_r+b_r+nrc=s_r+10*rc

然后通过枚举 da 来推算出 db 再判断是否合法即可,时间复杂度为 O(2^n n^2),你发现 dfs 的次数只跟进位有关系所以只有 O(4^{n/2})=O(2^n)。再根据附件所说的最小的回文数不会超过 10801 所以最多会算 207 次。 ::::info[code]

#include<bits/stdc++.h>
#define ll long long
using namespace std;
constexpr ll N=1e6+5,inf=1e18;
inline bool check(string s){
    if(s.empty()) return 0;
    int l=0,r=s.size()-1;
    while(l<r) if(s[l++]!=s[r--]) return 0;
    return 1;
}
string sub(string a,string b){
    string res;
    ll n=a.size(),m=b.size();
    ll jw=0;
    for(int i=0;i<n;i++){
        int da=a[n-1-i]-'0';
        int db=(i<m?b[m-1-i]-'0':0);
        int s=da-db-jw;
        if(s<0) s+=10,jw=1;
        else jw=0;
        res.push_back(s+'0');
    }
    while(res.size()>1 && res.back()=='0') res.pop_back();
    reverse(res.begin(),res.end());
    return res;
}
namespace two{
    int ls,la,lb;
    vector<ll> s,a,b;
    //lc 从 l-1 位向 l 的进位
    //rc 从 r 位向 r+1 位的进位
    bool dfs(int l,int lc,int rc){
        int r=la-l-1;
        if(l>r) return lc==rc;
        if(l==r){
            for(int da=(l==la-1?1:0);da<=9;da++){
                if(a[l]!=-1 && a[l]!=da) continue;
                for(int db=(l<lb?(l==lb-1?1:0):0);db<=(l<lb?9:0);db++){
                    if(l<lb && b[l]!=-1 && b[l]!=db) continue;
                    int sum=da+db+lc;
                    if(sum%10==s[l] && sum/10==rc){
                        a[l]=da;
                        if(l<lb) b[l]=db;
                        return 1;
                    }
                }
            }
            return 0;
        }
        //nlc 从 l 位向 l+1 位
        //nrc 从 r-1 位向 r 位
        for(int nlc=0;nlc<=1;nlc++){
            for(int nrc=0;nrc<=1;nrc++){
                //da+db+lc=s[l]+10*nlc
                int tl=s[l]+10*nlc-lc;//da+db
                if(tl<0 || tl>18) continue;
                //ad+db+nrc=s[r]+10*rc
                int tr=s[r]+10*rc-nrc;
                if(tr<0 || tr>18) continue;
                for(int da=(l==0?1:0);da<=9;da++){
                    if(a[l]!=-1 && a[l]!=da) continue;
                    if(a[r]!=-1 && a[r]!=da) continue;
                    int bl=tl-da;
                    if(l>=lb){
                        if(bl) continue;
                    }else{
                        int mnb=(l==lb-1?1:0);
                        if(bl<mnb || bl>9) continue;
                        if(b[l]!=-1 && b[l]!=bl) continue;
                        if(b[lb-1-l]!=-1 && b[lb-1-l]!=bl) continue;
                    }
                    int br=tr-da;
                    if(r>=lb){
                        if(br) continue;
                    }else{
                        int mnb=(r==lb-1?1:0);
                        if(br<mnb || br>9) continue;
                        if(b[r]!=-1 && b[r]!=br) continue;
                        if(b[lb-1-r]!=-1 && b[lb-1-r]!=br) continue;
                    }
                    if(l<lb && r<lb && lb-1-l==r && bl!=br) continue;
                    int old_al=a[l],old_ar=a[r];
                    int old_bl=(l<lb?b[l]:0),old_bll=(l<lb?b[lb-1-l]:0);
                    int old_br=(r<lb?b[r]:0),old_brr=(r<lb?b[lb-1-r]:0);
                    a[l]=a[r]=da;
                    if(l<lb) b[l]=b[lb-1-l]=bl;
                    if(r<lb) b[r]=b[lb-1-r]=br;
                    if(dfs(l+1,nlc,nrc)) return 1;
                    a[l]=old_al;a[r]=old_ar;
                    if(l<lb){b[l]=old_bl;b[lb-1-l]=old_bll;}
                    if(r<lb){b[r]=old_br;b[lb-1-r]=old_brr;}
                }
            }
        }
        return 0;
    }
    inline bool solve(string S,string& A,string& B){
        ls=S.size();s.resize(ls);
        for(int i=0;i<ls;i++) s[i]=S[ls-1-i]-'0';
        for(la=ls;la>=max(1,ls-1);la--){
            for(lb=la;lb>=1;lb--){
                a.assign(la,-1);
                b.assign(lb,-1);
                if(dfs(0,0,(la<ls?s[ls-1]:0))){
                    A.resize(la);
                    for(int i=0;i<la;i++) A[i]=a[la-1-i]+'0';
                    B.resize(lb);
                    for(int i=0;i<lb;i++) B[i]=b[lb-1-i]+'0';
                    return 1;
                }
            }
        }
        return 0;
    }
}
vector<string> zw;
inline void solve(){
    string s;cin>>s;
    if(check(s)){
        cout<<s<<"\n";
        return ;
    }
    string a,b;
    if(two::solve(s,a,b)){
        cout<<a<<" "<<b<<"\n";
        return ;
    }
    for(auto k:zw){
        string tmp=sub(s,k);
        if(check(tmp)){
            cout<<k<<" "<<tmp<<"\n";
            return ;
        }
        if(two::solve(tmp,a,b)){
            cout<<k<<" "<<a<<" "<<b<<"\n";
            return ;
        }
    }
}
signed main(){
    ios::sync_with_stdio(0),cin.tie(0),cout.tie(0);
    ll T;cin>>T;
    for(int i=1;i<=10801;i++){
        string t=to_string(i);
        if(check(t)) zw.push_back(t);
    }
    for(int _=1;_<=T;_++){
        cout<<"Case #"<<_<<": ";
        solve();
    }
    return 0;
}

::::