Official

H - Digit Circus Editorial by sheyasutaka


桁 DP と呼ばれる手法によって解くことができます.

DP 上では,扱う整数の桁数を \(N\) と揃えるため,leading-zero を特例として許すものとします.

扱う整数の冒頭 \(i\) 桁の状態を表す添字として,以下を持てば十分です.

  • 使う数字の集合 \(b\)
  • \(3\) で割ったあまり \(m\)
  • \(N\) の冒頭 \(i\) 桁と比べて小さいか等しいかを表す変数 \(l\)

状態 \(dp[i-1][b][m][l]\) および \(i+1\) 文字目の数字 \(d\) を固定したとき,これらの変数から遷移先の状態 \(dp[i][\cdot][\cdot][\cdot]\) の添字と係数が求まります.具体的な遷移は実装例を参照してください.

実装時の注意点として,leading-zero を置いたときに 0 を使ったことにならないようにする必要があります.

時間計算量は \(D=10\) とおいて \(O(2^D D \log N)\) です.

実装例 (C++)

#include <iostream>
using std::cin;
using std::cout;
using std::cerr;
using std::endl;
#include <vector>
using std::vector;
#include <string>
using std::string;
using std::to_string;
#include <bit>
using std::popcount;

typedef long long int ll;

const ll FOD = 998244353;

string s;

ll dp[505][1 << 10][3][2];
void solve() {
	// init dp[0]
	for (ll bi = 0; bi < (1<<10); bi++) {
		for (ll mi = 0; mi < 3; mi++) {
			for (ll li = 0; li < 2; li++) {
				dp[0][bi][mi][li] = 0;
			}
		}
	}
	dp[0][0][0][0] = 1;
	
	for (ll i = 0; i < s.size(); i++) {
		ll si = (s[i] - '0');

		// init dp[i+1]
		for (ll bi = 0; bi < (1<<10); bi++) {
			for (ll mi = 0; mi < 3; mi++) {
				for (ll li = 0; li < 2; li++) {
					dp[i+1][bi][mi][li] = 0;
				}
			}
		}

		// dp[i] -> dp[i+1]
		for (ll bi = 0; bi < (1<<10); bi++) {
			for (ll mi = 0; mi < 3; mi++) {
				// unroll loop for li
				for (ll d = 0; d <= 9; d++) {
					ll bx = ((bi == 0 && d == 0) ? 0 : (bi | (1LL << d)));
					ll mx = (mi * 10 + d) % 3;
					// li 0
					if (d <= si) {
						dp[i+1][bx][mx][(d == si) ? 0 : 1] += dp[i][bi][mi][0];
					}
					dp[i+1][bx][mx][1] += dp[i][bi][mi][1];
				}
			}
		}

		// modulo
		for (ll bi = 0; bi < (1<<10); bi++) {
			for (ll mi = 0; mi < 3; mi++) {
				for (ll li = 0; li < 2; li++) {
					dp[i+1][bi][mi][li] %= FOD;
				}
			}
		}
	}

	// dp[n] -> ans
	ll ans = 0;
	for (ll bi = 1; bi < (1<<10); bi++) {
		for (ll mi = 0; mi < 3; mi++) {
			for (ll li = 0; li < 2; li++) {
				ll cond = 0;
				if (popcount((uint64_t)bi) == 3) cond++;
				if (bi & (1LL << 3)) cond++;
				if (mi == 0) cond++;

				if (cond == 1) {
					ans += dp[s.size()][bi][mi][li];
				}
			}
		}
	}
	ans %= FOD;

	cout << ans << "\n";
}

int main (void) {
	std::cin.tie(nullptr);
	std::ios_base::sync_with_stdio(false);

	cin >> s;
	
	solve();

	return 0;
}

posted:
last update: