문제 소개
길이가 n인 두 개의 문자열 s1과 s2, 그리고 evil이라는 이름의 문자열 하나가 주어집니다. 목표는 '좋은(good) 문자열'의 개수를 구하는 것입니다.
어떤 문자열이 다음 조건을 모두 만족할 때 좋은 문자열이라고 정의합니다.
- 길이가 정확히 n이다.
- 사전 순으로 s1보다 크거나 같다.
- 사전 순으로 s2보다 작거나 같다.
- evil을 부분 문자열로 포함하지 않는다.
정답은 매우 커질 수 있으므로, 결과를 109 + 7로 나눈 나머지를 반환해야 합니다.
예시
n = 2, s1 = "bb", s2 = "db", evil = "a"가 입력으로 주어지면 출력은 51이 됩니다. 그 이유는 다음과 같습니다.
- 'b'로 시작하는 좋은 문자열 25개: "bb", "bc", "bd", ..., "bz"
- 'c'로 시작하는 좋은 문자열 25개: "cb", "cc", "cd", ..., "cz"
- 'd'로 시작하는 좋은 문자열 1개: "db"
즉 25 + 25 + 1 = 51개가 정답입니다.
풀이 접근 방법
이 문제는 자릿수 DP(digit DP) 기법과 KMP 알고리즘의 접두사 일치 개념을 결합하면 효율적으로 해결할 수 있습니다. 각 문자 위치에서 지금까지 만든 문자열이 상한 문자열 s와 일치하는지 여부(l 플래그)를 추적하고, 동시에 evil의 접두사와 일치하는 최대 길이를 상태로 관리합니다. 전이 테이블 tr[i][j]는 'evil의 접두사 i글자가 일치하는 상태에서 문자 j를 추가할 때, 새로 일치하게 되는 접두사 길이'를 저장합니다.
구체적인 풀이 단계는 다음과 같습니다.
- N := 500, M := 50으로 설정한다.
- 크기 (N+1) × (M+1) × 2의 배열 dp를 정의한다.
- 크기 (M+1) × 26의 배열 tr을 정의한다.
- m := 109 + 7로 설정한다.
- add() 함수를 정의한다. a, b를 받아 ((a mod m) + (b mod m)) mod m을 반환한다.
- solve() 함수를 정의한다. 이 함수는 n, s, e를 매개변수로 받으며 다음을 수행한다.
- 배열 e를 뒤집는다.
- tr과 dp를 0으로 채운다.
- i := 0부터 e의 크기 미만까지 1씩 증가시키며 반복한다.
- f := e의 인덱스 0부터 i−1까지의 부분 문자열
- j := 0부터 26 미만까지 1씩 증가시키며 반복한다.
- ns := f + ('a' + j에 해당하는 문자)
- k := i+1부터 1씩 감소시키며 반복한다.
- ns의 인덱스 (i+1−k)부터 끝까지의 부분 문자열이 e의 인덱스 0부터 k−1까지의 부분 문자열과 같으면 tr[i][j] := k로 설정하고 반복을 종료한다.
- m := e의 크기로 설정한다.
- i := 0부터 n 이하까지 1씩 증가시키며 반복한다.
- j := 0부터 m 미만까지 1씩 증가시키며 반복한다.
- dp[i][j][0] := 0, dp[i][j][1] := 0으로 초기화한다.
- j := 0부터 m 미만까지 1씩 증가시키며 반복한다.
- dp[n][0][1] := 1로 설정한다.
- i := n−1부터 0 이상일 때까지 1씩 감소시키며 반복한다.
- j := 0부터 e의 크기 미만까지 1씩 증가시키며 반복한다.
- k := 0부터 26 미만까지 1씩 증가시키며 반복한다.
- l을 {0, 1} 범위에서 순회한다.
- k > s[i] − 'a'이면 nl := 0
- k < s[i] − 'a'이면 nl := 1
- 그 외의 경우 nl := l
- dp[i][tr[j][k]][nl] := add(dp[i][tr[j][k]][nl], dp[i+1][j][l])
- l을 {0, 1} 범위에서 순회한다.
- k := 0부터 26 미만까지 1씩 증가시키며 반복한다.
- j := 0부터 e의 크기 미만까지 1씩 증가시키며 반복한다.
- ret := 0으로 설정한다.
- i := 0부터 e의 크기 미만까지 1씩 증가시키며 반복한다.
- ret := add(ret, dp[0][i][1])
- ret을 반환한다.
메인 메서드 처리 과정
- ok := 1로 초기화한다.
- i := 0부터 s1의 크기 미만이고 ok가 참인 동안 1씩 증가시키며 반복한다.
- ok := (s1[i] == 'a')
- ok가 거짓이면 다음을 수행한다.
- i := s1의 크기 − 1부터 0 이상일 때까지 1씩 감소시키며 반복한다.
- s1[i]가 'a'가 아니면 s1[i]를 1 감소시키고 반복을 종료한다.
- 그렇지 않으면 s1[i] := 'z'로 설정한다.
- i := s1의 크기 − 1부터 0 이상일 때까지 1씩 감소시키며 반복한다.
- left := (ok가 참이면 0, 아니면 solve(n, s1, evil))
- right := solve(n, s2, evil)
- (right − left + m) mod m을 반환한다.
여기서 left는 "s1 − 1" 이하인 문자열의 개수, right는 "s2" 이하인 문자열의 개수를 의미합니다. s1을 실제로 1 감소시켜 처리함으로써 's1 미만'의 개수를 구하고, right에서 left를 빼면 [s1, s2] 범위에 속하는 좋은 문자열의 개수를 얻을 수 있습니다.
구현 예제
아래 코드를 통해 더 잘 이해해 보겠습니다.
#include <bits/stdc++.h>
using namespace std;
typedef long long int lli;
const int N = 500;
const int M = 50;
int dp[N + 1][M + 1][2];
int tr[M + 1][26];
const lli m = 1e9 + 7;
class Solution {
public:
int add(lli a, lli b){
return ((a % m) + (b % m)) % m;
}
lli solve(int n, string s, string e){
reverse(e.begin(), e.end());
memset(tr, 0, sizeof(tr));
memset(dp, 0, sizeof(dp));
for (int i = 0; i < e.size(); i++) {
string f = e.substr(0, i);
for (int j = 0; j < 26; j++) {
string ns = f + (char)(j + 'a');
for (int k = i + 1;; k--) {
if (ns.substr(i + 1 - k) == e.substr(0, k)) {
tr[i][j] = k;
break;
}
}
}
}
int m = e.size();
for (int i = 0; i <= n; i++) {
for (int j = 0; j < m; j++) {
dp[i][j][0] = dp[i][j][1] = 0;
}
}
dp[n][0][1] = 1;
for (int i = n - 1; i >= 0; i--) {
for (int j = 0; j < e.size(); j++) {
for (int k = 0; k < 26; k++) {
for (int l : { 0, 1 }) {
int nl;
if (k > s[i] - 'a') {
nl = 0;
}
else if (k < s[i] - 'a') {
nl = 1;
}
else
nl = l;
dp[i][tr[j][k]][nl] = add(dp[i][tr[j][k]]
[nl], dp[i + 1][j][l]);
}
}
}
}
lli ret = 0;
for (int i = 0; i < e.size(); i++) {
ret = add(ret, dp[0][i][1]);
}
return ret;
}
int findGoodStrings(int n, string s1, string s2, string evil) {
bool ok = 1;
for (int i = 0; i < s1.size() && ok; i++) {
ok = s1[i] == 'a';
}
if (!ok) {
for (int i = s1.size() - 1; i >= 0; i--) {
if (s1[i] != 'a') {
s1[i]--;
break;
}
s1[i] = 'z';
}
}
int left = ok ? 0 : solve(n, s1, evil);
int right = solve(n, s2, evil);
return (right - left + m) % m;
}
};
main(){
Solution ob;
cout << (ob.findGoodStrings(2, "bb", "db", "a"));
}
입력
2, "bb", "db", "a"
출력
51