문제 이해하기
여러 개의 상자가 일렬로 나열되어 있고, 각 상자는 서로 다른 색상을 가지고 있다고 가정해 보겠습니다. 색상은 서로 다른 양의 정수로 표현됩니다. 우리는 모든 상자가 사라질 때까지 여러 라운드에 걸쳐 상자를 제거할 수 있으며, 각 라운드마다 같은 색상으로 연속된 k개의 상자(k ≥ 1)를 선택해 한 번에 제거하고 그 대가로 k × k점을 얻습니다.
예를 들어 입력이 [1, 3, 2, 2, 2, 4, 4, 3, 1]이라면 출력은 21이 됩니다. 목표는 상자를 제거하는 순서를 잘 선택해서 얻을 수 있는 최대 점수를 구하는 것입니다.
접근 방식: 구간 DP와 메모이제이션
이 문제는 단순한 탐욕(greedy) 방식으로는 최적해를 보장할 수 없습니다. 상자를 제거하는 순서에 따라 나중에 합쳐질 수 있는 같은 색상 그룹의 크기가 달라지기 때문입니다. 따라서 구간 기반의 동적 계획법(Dynamic Programming)과 메모이제이션을 활용해야 합니다.
핵심 아이디어는 구간 [i, j]를 처리할 때, 구간 왼쪽 바깥에 이미 boxes[i]와 같은 색상의 상자 k개가 붙어 있어서 나중에 합쳐질 수 있다는 정보까지 함께 고려하는 것입니다. 이를 위해 3차원 DP 배열 dp[i][j][k]를 사용합니다.
구체적인 해결 단계는 다음과 같습니다.
- solve() 함수를 정의합니다. 이 함수는 배열 boxes, 인덱스 i, j, 개수 k, 그리고 3차원 배열 dp를 매개변수로 받습니다.
- i > j이면 0을 반환합니다.
- dp[i][j][k]의 값이 -1이 아니라면(이미 계산된 값이라면) 해당 값을 그대로 반환합니다.
- ret := -∞(매우 작은 값)으로 초기화합니다.
- i + 1 ≤ j이고 boxes[i + 1]이 boxes[i]와 같은 동안 i와 k를 1씩 증가시켜, 시작 지점의 연속된 같은 색상 상자를 하나의 그룹으로 묶습니다.
- ret := max(ret, (k + 1) × (k + 1) + solve(boxes, i + 1, j, 0, dp))로 갱신합니다. 즉, 현재 그룹을 즉시 제거하는 경우를 계산합니다.
- x := i + 1부터 x ≤ j까지 반복하면서 다음을 수행합니다.
- boxes[x]가 boxes[i]와 같다면, 중간 구간 [i + 1, x − 1]을 먼저 모두 제거한 뒤 뒤쪽의 같은 색상 상자들을 앞 그룹과 합쳐서 제거하는 경우를 계산합니다. 즉, ret := max(ret, solve(boxes, i + 1, x − 1, 0, dp) + solve(boxes, x, j, k + 1, dp))로 갱신합니다.
- 마지막으로 dp[i][j][k] = ret을 저장하고 반환합니다.
메인 함수에서는 다음과 같이 처리합니다.
- n := boxes 배열의 크기
- (n + 1) × (n + 1) × (n + 1) 크기의 3차원 배열 dp를 선언하고 모든 값을 -1로 초기화합니다.
- solve(boxes, 0, n − 1, 0, dp)의 결과를 반환합니다.
C++ 구현 예제
#include <bits/stdc++.h>
using namespace std;
class Solution {
public:
int solve(vector <int>& boxes, int i, int j, int k, vector < vector < vector <int > > >& dp){
if(i > j) return 0;
if(dp[i][j][k] != -1) return dp[i][j][k];
int ret = INT_MIN;
for(; i + 1 <= j && boxes[i + 1] == boxes[i]; i++, k++);
ret = max(ret, (k + 1) * (k + 1) + solve(boxes, i + 1, j, 0, dp));
for(int x = i + 1; x <= j; x++){
if(boxes[x] == boxes[i]){
ret = max(ret, solve(boxes, i + 1, x - 1, 0, dp) + solve(boxes, x, j, k + 1, dp));
}
}
return dp[i][j][k] = ret;
}
int removeBoxes(vector<int>& boxes) {
int n = boxes.size();
vector < vector < vector <int > > > dp(n + 1, vector < vector <int> > (n + 1, vector <int>(n + 1, -1)));
return solve(boxes, 0, n - 1, 0, dp);
}
};
main(){
Solution ob;
vector<int> v = {1,3,2,2,2,4,4,3,1};
cout << (ob.removeBoxes(v));
}입력
{1,3,2,2,2,4,4,3,1}출력
21