数位dp 学习笔记

在博客园食用更佳

结合题目P2657 [SCOI2009] windy 数P13085 [SCOI2009] windy 数(加强版)进行讲解。

概念

引用oi-wiki中的一段话定义 数位dp:

数位dp:用来解决一类特定问题,这种问题比较好辨认,一般具有这几个特征:

  1. 要求统计满足一定条件的数的数量(即,最终目的为计数);
  2. 这些条件经过转化后可以使用「数位」的思想去理解和判断;
  3. 输入会提供一个数字区间(有时也只提供上界)来作为统计的限制;
  4. 上界很大(比如 \(10^{18}\)),暴力枚举验证会超时。

方法

解决 数位dp 题目的方法可以归结为五个字:拆、搜、记、限、零。

“拆”是指把数字拆开。本题中数据较小,但是遇到 \(10^{200}\) 这类问题,就需要把输入的字符串一位一位地给拆进数组里。

为什么这样做?因为 数位dp,顾名思义,是一种基于数位的 dp(废话),它的实现需要用到数字的每一位,不拆成数组会很麻烦。

这里的“搜”指搜索,大部分 数位dp 都可以用搜索实现。

“记”是指记忆化,做 dp 不用记忆化,只要数据是认真造的,一定会超时,复杂度爆炸。所以记忆化对身体复杂度很友好。

“限”指上限,在代码中一般体现为搜索中的参数 limit。这个参数负责记录你现在是否贴着上限。在 limittrue 的情况下,每一次递归都需要特殊判断,防止其超出上限。

“零”指前导零。

在 数位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;
}
posted @ 2026-07-20 15:56  cath20  阅读(7)  评论(0)    收藏  举报