数位dp 学习笔记
结合题目P2657 [SCOI2009] windy 数与P13085 [SCOI2009] windy 数(加强版)进行讲解。
概念
引用oi-wiki中的一段话定义 数位dp:
数位dp:用来解决一类特定问题,这种问题比较好辨认,一般具有这几个特征:
- 要求统计满足一定条件的数的数量(即,最终目的为计数);
- 这些条件经过转化后可以使用「数位」的思想去理解和判断;
- 输入会提供一个数字区间(有时也只提供上界)来作为统计的限制;
- 上界很大(比如 \(10^{18}\)),暴力枚举验证会超时。
方法
解决 数位dp 题目的方法可以归结为五个字:拆、搜、记、限、零。
拆
“拆”是指把数字拆开。本题中数据较小,但是遇到 \(10^{200}\) 这类问题,就需要把输入的字符串一位一位地给拆进数组里。
为什么这样做?因为 数位dp,顾名思义,是一种基于数位的 dp(废话),它的实现需要用到数字的每一位,不拆成数组会很麻烦。
搜
这里的“搜”指搜索,大部分 数位dp 都可以用搜索实现。
记
“记”是指记忆化,做 dp 不用记忆化,只要数据是认真造的,一定会超时,复杂度爆炸。所以记忆化对身体复杂度很友好。
限
“限”指上限,在代码中一般体现为搜索中的参数 limit。这个参数负责记录你现在是否贴着上限。在 limit 为 true 的情况下,每一次递归都需要特殊判断,防止其超出上限。
零
“零”指前导零。
在 数位dp 中,有一个板子题让你统计某个范围内所有数字出现的次数。在这里我们很容易对前导零进行统计,所以应该有一个参数 lead 用以记录现在是否是前导零。对于是前导零,即 lead 的值为 true 的情况,也应该特殊处理。
模板
提供一个 数位dp 的代码模板:
#include<bits/stdc++.h>
#define int long long
using namespace std;
int len,a[1000],f[1000][2][2];
string n,m;
int dfs(int now,bool lead,bool limit)
{
if(now==len+1) return ???;
if(f[now][lead][limit]!=-1) return f[now][lead][limit];
int sum=0,up=???;
for(int i=0;i<=up;++i) sum+=dfs(now+1,(i==0&&lead),(i==a[now]&&limit));
return f[now][lead][limit]=sum;
}
int solve(string s,int z)
{
if((s=="1"&&z)||s=="0") return 1;
len=s.size();
memset(f,-1,sizeof(f));
for(int i=1;i<=len;++i) a[i]=s[i-1]-'0';
a[len]-=z;
for(int i=len;i>=1&&a[i]==-1;--i)
{
a[i]=9;
--a[i-1];
}
if(a[1]==0)
{
--len;
for(int i=1;i<=len;++i) a[i]=a[i+1];
}
return dfs(1,true,true);
}
signed main()
{
ios::sync_with_stdio(false);
cin.tie();
cout.tie();
cin>>n>>m;
cout<<solve(m,0)-solve(n,1);
return 0;
}
例题
之所以选择这道题作为例题,是因为它完美涵盖了这五个字,并且足够板子足够简单。
所以对于这道题,我们的思路如下。
首先,题目要求计算区间 \(a\) 到 \(b\) 之间的所有 windy数。那么我们是不是应该再加一个参数用以记录下限呢?不,我们可以使用前缀和的思想,用 \(1\) ~ \(a\) 的结果减去 \(1\) ~ \(b-1\) 的结果。
然后,我们就可以开始填板子了。
#include<bits/stdc++.h>
#define int long long
using namespace std;
int len,a[1000],f[1000][2][2][12];
string n,m;
int dfs(int now,bool lead,bool limit,int last)
{
if(now==len+1) return 1;
if(f[now][lead][limit][last]!=-1) return f[now][lead][limit][last];
int sum=0,up=9;
if(limit) up=a[now];
for(int i=0;i<=up;++i) if(abs(i-last)>=2||lead) sum+=dfs(now+1,(i==0&&lead),(i==a[now]&&limit),i);
return f[now][lead][limit][last]=sum;
}
int solve(string s,int z)
{
if((s=="1"&&z)||s=="0") return 1;
len=s.size();
memset(f,-1,sizeof(f));
for(int i=1;i<=len;++i) a[i]=s[i-1]-'0';
a[len]-=z;
for(int i=len;i>=1&&a[i]==-1;--i)
{
a[i]=9;
--a[i-1];
}
if(a[1]==0)
{
--len;
for(int i=1;i<=len;++i) a[i]=a[i+1];
}
return dfs(1,true,true,11);
}
signed main()
{
ios::sync_with_stdio(false);
cin.tie();
cout.tie();
cin>>n>>m;
cout<<solve(m,0)-solve(n,1);
return 0;
}

浙公网安备 33010602011771号