ZR 集训 Day 9 – 状态压缩 DP,数位 DP

ooliver 发布于 16 小时前 82 次阅读 OI


AI 摘要

从吃奶酪的最短路径,到省选滚榜的排列玄机,再到萌数判重与二进制下的对数求和——状态压缩与数位 DP 总能用精巧的状态设计,把看似无从下手的计数与最优化问题一击必杀。

状态压缩 DP,数位 DP

P1433 吃奶酪

设 $dp_{i,j}$ 表示当前选择的集合为 $i$ 时,最后一个拿到的奶酪为 $k$,所走的最小距离,就可以枚举最后一个点和走到的新点进行转移,复杂度 $O(2^nn^2)$。

代码:

C++
#include<bits/stdc++.h>
using namespace std;

int n;
double ans=1e9;
double x[25],y[25],dp[(1<<15)+5][25];

double dis(double x1,double y1,double x2,double y2){
    return sqrt((x1-x2)*(x1-x2)+(y1-y2)*(y1-y2));
}

int main(){
    cin>>n;
    for(int i=1;i<=n;i++) cin>>x[i]>>y[i];
    memset(dp,127,sizeof dp);
    for(int i=1;i<=n;i++) dp[1<<(i-1)][i]=dis(0,0,x[i],y[i]);
    for(int i=0;i<(1<<n);i++)
        for(int j=1;j<=n;j++) if((i>>(j-1))&1)
            for(int k=1;k<=n;k++) if((i>>(k-1))&1){
                dp[i][k]=min(dp[i][k],dp[i-(1<<(k-1))][j]+dis(x[k],y[k],x[j],y[j]));
            }
    for(int i=1;i<=n;i++) ans=min(ans,dp[(1<<n)-1][i]);
    printf("%0.2f",ans);
    return 0;
}

P7519 [省选联考 2021 A/B 卷] 滚榜

题目要求的是排列的方案数,所以我们不需要求出具体的 $b_i$,考虑记录其差分数组,每次加上对后面的影响,并且转移时贪心地按照最优策略转移。

代码:

C++
#include<bits/stdc++.h>
using namespace std;

#define int long long
int n,m,x,id,ans;
int a[15],dp[(1<<13)+5][15][505];

signed main(){
    cin>>n>>m;
    for(int i=1;i<=n;i++){
        cin>>a[i];
        if(a[i]>x) x=a[i],id=i;
    }
    dp[0][0][0]=1;
    for(int i=1;i<=n;i++) if(n*max(a[id]-a[i]+(id<i),0ll)<=m) dp[1<<(i-1)][i][n*max(a[id]-a[i]+(id<i),0ll)]=1;
    for(int s=1;s<(1<<n);s++)
        for(int i=0;i<n;i++) if(~(s>>i)&1)
            for(int j=0;j<n;j++) if((s>>j)&1)
                for(int v=max(0ll,a[j+1]-a[i+1]+(i>j))*(n-__builtin_popcount(s));v<=m;v++){
                    dp[s|(1<<i)][i+1][v]+=dp[s][j+1][v-max(0ll,a[j+1]-a[i+1]+(i>j))*(n-__builtin_popcount(s))];
                }
    for(int i=0;i<=m;i++) for(int j=1;j<=n;j++) ans+=dp[(1<<n)-1][j][i];
    cout<<ans;
    return 0;
}

P4127 [AHOI2009] 同类分布

数位 DP 模板题,先枚举模数,

设 $dp_{p,s,m,y}$ 表示当前数长度为 $p$、数位和为 $s$、数字取模为 $m$,$y$ 表示有没有上界限制,这样情况下的答案,直接记忆化搜索。

代码:

C++
#include<bits/stdc++.h>
using namespace std;

#define int long long
const int N=20;
int n,len,mod;
int a[N],dp[N][200][200][2];

int dfs(int p,int s,int m,int y){
    if(p>len) return (s==mod&&m==0);
    if(dp[p][s][m][y]!=-1) return dp[p][s][m][y];
    int mx=y?a[len-p+1]:9,res=0;
    for(int i=0;i<=mx;i++) res+=dfs(p+1,s+i,(m*10+i)%mod,i==mx&&y);
    return dp[p][s][m][y]=res;
}

int solve(int x){
    int res=0;
    len=0;
    while(x) a[++len]=x%10,x/=10;
    for(mod=1;mod<=9*len;mod++){
        memset(dp,-1,sizeof dp);
        res+=dfs(1,0,0,1);
    }
    return res;
}

signed main(){
    int l,r;
    cin>>l>>r;
    cout<<solve(r)-solve(l-1);
    return 0;
}

P3107 [USACO14OPEN] Odometer S

比较套路,枚举出现次数最多的那个数进行记忆化搜索即可。

注意会出现一个特例,就是一个数中有两个数字出现次数一样多,再做一遍记忆化把它容斥掉即可。

代码:

C++
#include<bits/stdc++.h>
using namespace std;

#define int long long
const int N=20;
int len;
int a[N],dp[N][N][N][2][2];

int dfs1(int p,int k,int s1,int s2,int t1,int t2){
    if(p>len) return (s1>0&&s1>=s2);
    if(dp[p][s1][s2][t1][t2]!=-1) return dp[p][s1][s2][t1][t2];
    int mx=t1?a[len-p+1]:9,res=0;
    for(int i=0;i<=mx;i++) res+=dfs1(p+1,k,s1+(i!=0||t2==0)*(i==k),s2+(i!=0||t2==0)*(i!=k),i==mx&&t1,i==0&&t2);
    return dp[p][s1][s2][t1][t2]=res;
}

int dfs2(int p,int x,int y,int s1,int s2,int t1,int t2){
    if(p>len) return (s1>0&&s1==s2);
    if(dp[p][s1][s2][t1][t2]!=-1) return dp[p][s1][s2][t1][t2];
    int mx=t1?a[len-p+1]:9,res=0;
    if(x<=mx||t1==0) res+=dfs2(p+1,x,y,s1+(x!=0||t2==0),s2,x==mx&&t1,x==0&&t2);
	if(y<=mx||t1==0) res+=dfs2(p+1,x,y,s1,s2+(y!=0||t2==0),y==mx&&t1,y==0&&t2);
	if(x!=0&&y!=0&&t2==1) res+=dfs2(p+1,x,y,s1,s2,mx==0&&t1,t2);
    return dp[p][s1][s2][t1][t2]=res;
}

int s1(int num,int k){
    len=0;
    while(num) a[++len]=num%10,num/=10;
    memset(dp,-1,sizeof dp);
    return dfs1(1,k,0,0,1,1);
}

int s2(int num,int x,int y){
    len=0;
    while(num) a[++len]=num%10,num/=10;
    memset(dp,-1,sizeof dp);
    return dfs2(1,x,y,0,0,1,1);
}

signed main(){
    int l,r,ans=0;
    cin>>l>>r;
    for(int i=0;i<=9;i++) ans+=s1(r,i)-s1(l-1,i);
    for(int i=0;i<=9;i++) for(int j=i+1;j<=9;j++) ans-=s2(r,i,j)-s2(l-1,i,j);
    cout<<ans;
    return 0;
}

P9821 [ICPC 2020 Shanghai R] Sum of Log

依旧数位 DP,感觉也挺套路的,只不过是换成了二进制下的,其他好像没啥好说的了。

代码:

C++
#include<bits/stdc++.h>
using namespace std;

#define int long long
const int N=32,mod=1e9+7;
int n,m;
int dp[N][2][2][2];

int dfs(int p,int t1,int t2,int t3){
    if(p<0) return 1;
    if(dp[p][t1][t2][t3]!=-1) return dp[p][t1][t2][t3];
    int m1=t1?(n>>p)&1:1,m2=t2?(m>>p)&1:1,res=0;
    for(int i=0;i<=m1;i++) for(int j=0;j<=m2;j++){
        if(i&j) continue;
        (res+=((t3&&(i||j))?(p+1):1)*dfs(p-1,(i==m1)&&t1,(j==m2)&&t2,t3&&!i&&!j))%=mod;
    }
    return (dp[p][t1][t2][t3]=res)%mod;
}

int solve(int x,int y){
    n=x,m=y;
    memset(dp,-1,sizeof dp);
    return (dfs(30,1,1,1)-1+mod)%mod;
}

signed main(){
    int T,l,r;
    cin>>T;
    while(T--){
        cin>>l>>r;
        cout<<(solve(l,r))%mod<<"\n";
    }
    return 0;
}

P3413 SAC#1 - 萌数

需要想到一个点:长度大于或等于 $2$ 的子串,其实我们只用判断两种情况:AA 和 ABA。

其余没什么特别的,注意 $n$ 的范围很大。

代码:

C++
#include<bits/stdc++.h>
using namespace std;

#define int long long
const int N=1005,mod=1000000007;
string l,r;
int n,a[N],dp[N][11][11][2];

int dfs(int pos,int l1,int l2,bool has,bool lim,bool lead){
    if(pos>n) return has;
    if(!lim&&!lead&&dp[pos][l1][l2][has]!=-1) return dp[pos][l1][l2][has];
    int mx=lim?a[pos]:9,res=0;
    for(int i=0;i<=mx;i++){
        if(lead&&!i) res=(res+dfs(pos+1,10,10,0,lim&&i==mx,1))%mod;
        else{
            bool nh=has;
            if(!lead){
                if(l1!=10&&i==l1) nh=1;
                if(l2!=10&&i==l2) nh=1;
            }
            res=(res+dfs(pos+1,i,l1,nh,lim&&i==mx,0))%mod;
        }
    }
    if(!lim&&!lead) dp[pos][l1][l2][has]=res;
    return res;
}

int solve(string s){
    n=s.size();
    for(int i=1;i<=n;i++) a[i]=s[i-1]-'0';
    memset(dp,-1,sizeof(dp));
    return dfs(1,10,10,0,1,1);
}

bool check(string s){
    for(int i=0;i<s.size();i++){
        if(i+1<s.size()&&s[i]==s[i+1]) return 1;
        if(i+2<s.size()&&s[i]==s[i+2]) return 1;
    }
    return 0;
}

signed main(){
    cin>>l>>r;
    int ans=(solve(r)-solve(l)+mod)%mod;
    if(check(l)) ans=(ans+1)%mod;
    cout<<ans<<"\n";
    return 0;
}