记录位置, 当前数对m取模的结果以及当前位是否为偶数位. 没想到数位DP居然TLE了, 求助Orz 代码如下:
#include <bits/stdc++.h>
#define debug() freopen("../test.in", "r", stdin)
#define max(a, b) ((a) > (b) ? (a) : (b))
#define min(a, b) ((a) < (b) ? (a) : (b))
#define abs(x) ((x) >= 0 ? (x) : -(x))
#define mst(x, y) memset((x), (y), sizeof (x))
#define endl '\n'
using namespace std;
__attribute__((unused)) typedef long long ll;
__attribute__((unused)) typedef unsigned long long ull;
__attribute__((unused)) inline ll read() {
ll x = 0, f = 1;
char ch = getchar();
while (ch < '0' || ch > '9') {
if (ch == '-')
f = -1;
ch = getchar();
}
while (ch >= '0' && ch <= '9') {
x = (x << 1) + (x << 3) + (ch ^ 48);
ch = getchar();
}
return x * f;
}
__attribute__((unused)) const long double Pi = 3.1415926535897932384626433832795;
__attribute__((unused)) const ll maxN = 2e3 + 4, maxM = 1e2 + 4, maxA = 14, maxB = 34, INF = 0x3f3f3f3f3f3f3f3f, MOD = 1e9 + 7;
ll m, d;
string l, r;
int len;
ll dp[maxN][maxN][2] = {};
int a[maxN] = {};
ll dfs(int pos, int sum, bool isEven, bool lead, bool limit) {
if (pos > len)
return !sum;
if (dp[pos][sum][isEven] != -1 && !lead && !limit)
return dp[pos][sum][isEven] % MOD;
ll ans = 0;
int res = limit ? a[pos] : 9;
for (int i = 0; i <= res; ++i) {
if (!(lead && !i) && (isEven && i != d || !isEven && i == d))
continue;
bool li = limit && i == res;
if (lead && !i) {
ans = (ans + dfs(pos + 1, 0, false, true, li) % MOD) % MOD;
continue;
}
ans = (ans + dfs(pos + 1, (sum * 10 + i) % m, !isEven, false, li) % MOD) % MOD;
}
if (!lead && !limit)
dp[pos][sum][isEven] = ans % MOD;
return ans % MOD;
}
ll solve(const string &num) {
len = 0;
for (auto i : num) {
a[++len] = i - '0';
}
for (int i = 1; i <= len; ++i) {
for (auto & j : dp[i]) {
for (auto & k : j) {
k = -1;
}
}
}
return dfs(1, 0, false, true, true) % MOD;
}
int check(const string &num) {
bool isEven = false;
int sum = 0;
for (auto i : num) {
if (isEven && i - '0' != d || !isEven && i - '0' == d)
return 0;
sum = (sum * 10 + i - '0') % m;
isEven = !isEven;
}
return !sum;
}
int main() {
// debug();
m = read(), d = read();
cin >> l >> r;
if (l == r) {
cout << check(l);
return 0;
}
cout << ((solve(r) - solve(l)) % MOD + check(l) + MOD) % MOD;
}