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

C++에서 가장 짧은 슈퍼스트링 찾기: 비트마스크 DP 완벽 가이드

문자열 배열 A가 주어졌을 때, A에 속한 모든 문자열을 부분 문자열(substring)로 포함하는 가장 짧은 문자열, 즉 '최단 슈퍼스트링(shortest superstring)'을 찾아야 합니다. 이때 배열 A 안의 어떤 문자열도 다른 문자열의 부분 문자열이 아니라고 가정할 수 있습니다.

예를 들어 입력이 ["dbsh", "dsbbhs", "hdsb", "ssdb", "bshdbsd"]라면 출력은 "hdsbbhssdbshdbsd"가 됩니다.

문제 해결 접근 방식

이 문제는 외판원 순환(TSP) 문제와 구조가 유사하며, 비트마스크(bitmask)와 동적 계획법(DP)을 조합하면 효율적으로 해결할 수 있습니다. 핵심 아이디어는 두 문자열을 이어 붙일 때 서로 겹치는 구간을 최대한 재활용하여 전체 길이를 줄이는 것입니다.

1단계: 겹침 길이 계산 함수 calc()

calc(a, b)는 문자열 a 뒤에 b를 이어 붙일 때 새로 추가해야 하는 문자 수를 반환합니다.

  • i를 0부터 a의 길이 미만까지 1씩 증가시키며 반복합니다.
  • a의 i번째 인덱스부터 끝까지 잘라낸 부분 문자열이 b의 맨 앞과 일치하면, b.size() - a.size() + i를 반환합니다.
  • 반복이 끝날 때까지 겹치는 구간을 찾지 못했다면 b의 전체 길이를 반환합니다.

2단계: 비용 행렬(graph) 구성

  • 결괏값을 담을 문자열 ret을 빈 문자열로 초기화하고, n := A의 크기로 설정합니다.
  • n × n 크기의 2차원 배열 graph를 만들고, graph[i][j] := calc(A[i], A[j]), graph[j][i] := calc(A[j], A[i])를 저장합니다. 이 값은 'i번 문자열 다음에 j번 문자열을 붙일 때 추가되는 길이'를 의미합니다.

3단계: 비트마스크 DP 테이블 채우기

  • 크기가 2n × n인 dp 배열과 path 배열을 선언합니다.
  • minVal := 무한대(inf), last := -1로 초기화합니다.
  • 모든 dp[i][j] 값을 무한대로 초기화합니다.
  • i를 0부터 2n 미만까지, j를 0부터 n 미만까지 순회하며 다음을 수행합니다.
    • i AND 2j 값이 0이 아니면(즉, j번째 문자열이 현재 집합 i에 포함되어 있으면) prev := i XOR 2j를 계산합니다.
    • prev가 0이면(j번 문자열 하나만 선택한 상태) dp[i][j] := A[j]의 길이로 설정합니다.
    • 그렇지 않으면 k를 0부터 n 미만까지 순회하며, prev AND 2k가 참이고 dp[prev][k]가 무한대가 아니며 dp[prev][k] + graph[k][j] < dp[i][j]를 만족할 때 dp[i][j] := dp[prev][k] + graph[k][j], path[i][j] := k로 갱신합니다.
  • i가 2n − 1(모든 문자열을 사용한 상태)이고 dp[i][j] < minVal이면 minVal := dp[i][j], last := j로 갱신합니다.

4단계: 경로 복원 및 결과 문자열 생성

  • curr := 2n − 1로 설정하고 스택 st를 생성합니다.
  • curr > 0인 동안 last를 스택에 push하고, temp := curr, curr := curr − 2last, last := path[temp][last]를 반복합니다.
  • 스택에서 첫 번째 원소를 꺼내 i에 저장하고 ret에 A[i]를 더합니다.
  • 스택이 빌 때까지 j := 스택 top, pop을 수행한 뒤, ret에 A[j]의 마지막 graph[i][j]개 문자(부분 문자열)를 이어 붙이고 i := j로 갱신합니다.
  • 최종적으로 ret을 반환합니다.

이 알고리즘의 시간 복잡도는 O(2n × n2)로, 문자열 개수 n이 작을 때(대략 20 이하) 실용적으로 동작합니다.

C++ 구현 예제

#include <bits/stdc++.h>
using namespace std;
class Solution {
   public:
   int calc(string& a, string& b){
      for (int i = 0; i < a.size(); i++) {
         if (b.find(a.substr(i)) == 0) {
            return b.size() - a.size() + i;
         }
      }
      return (int)b.size();
   }
   string shortestSuperstring(vector<string>& A){
      string ret = "";
      int n = A.size();
      vector<vector<int> > graph(n, vector<int>(n));
      for (int i = 0; i < n; i++) {
         for (int j = 0; j < n; j++) {
            graph[i][j] = calc(A[i], A[j]);
            graph[j][i] = calc(A[j], A[i]);
         }
      }
      int dp[1 << n][n];
      int path[1 << n][n];
      int minVal = INT_MAX;
      int last = -1;
      for (int i = 0; i < (1 << n); i++)
      for (int j = 0; j < n; j++)
      dp[i][j] = INT_MAX;
      for (int i = 1; i < (1 << n); i++) {
         for (int j = 0; j < n; j++) {
            if ((i & (1 << j))) {
               int prev = i ^ (1 << j);
               if (prev == 0) {
                  dp[i][j] = A[j].size();
               } else {
                  for (int k = 0; k < n; k++) {
                     if ((prev & (1 << k)) && dp[prev][k] !=
                     INT_MAX && dp[prev][k] + graph[k][j] < dp[i][j]) {
                        dp[i][j] = dp[prev][k] + graph[k][j];
                        path[i][j] = k;
                     }
                  }
               }
            }
            if (i == (1 << n) - 1 && dp[i][j] < minVal) {
               minVal = dp[i][j];
               last = j;
            }
         }
      }
      int curr = (1 << n) - 1;
      stack<int> st;
      while (curr > 0) {
         st.push(last);
         int temp = curr;
         curr -= (1 << last);
         last = path[temp][last];
      }
      int i = st.top();
      st.pop();
      ret += A[i];
      while (!st.empty()) {
         int j = st.top();
         st.pop();
         ret += (A[j].substr(A[j].size() - graph[i][j]));
         i = j;
      }
      return ret;
   }
};
main(){
   Solution ob;
   vector<string> v = {"dbsh","dsbbhs","hdsb","ssdb","bshdbsd"};
   cout << (ob.shortestSuperstring(v));
}

입력

{"dbsh","dsbbhs","hdsb","ssdb","bshdbsd"}

출력

hdsbbhssdbshdbsd