Computer >> 컴퓨터 >  >> 프로그래밍 >> C++

C++에서 모든 '좋은 문자열' 찾기 – 동적 계획법(DP) 풀이


문제 소개

길이가 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으로 초기화한다.
  • 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])
  • 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'로 설정한다.
  • 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