ALGORITHM NOTE2

BOJ 11049 - 행렬 곱셈 순서

어째서 곱하는 순서도 따져야 되는거야!

#algorithm#boj#gold#dp
아카이브로 돌아가기

문제 링크

문제

크기가 N×M인 행렬 A와 M×K인 B를 곱할 때 필요한 곱셈 연산의 수는 총 N×M×K번이다. 행렬 N개를 곱하는데 필요한 곱셈 연산의 수는 행렬을 곱하는 순서에 따라 달라지게 된다.

예를 들어, A의 크기가 5×3이고, B의 크기가 3×2, C의 크기가 2×6인 경우에 행렬의 곱 ABC를 구하는 경우를 생각해보자.

  • AB를 먼저 곱하고 C를 곱하는 경우 (AB)C에 필요한 곱셈 연산의 수는 5×3×2 + 5×2×6 = 30 + 60 = 90번이다.
  • BC를 먼저 곱하고 A를 곱하는 경우 A(BC)에 필요한 곱셈 연산의 수는 3×2×6 + 5×3×6 = 36 + 90 = 126번이다.

같은 곱셈이지만, 곱셈을 하는 순서에 따라서 곱셈 연산의 수가 달라진다.

행렬 N개의 크기가 주어졌을 때, 모든 행렬을 곱하는데 필요한 곱셈 연산 횟수의 최솟값을 구하는 프로그램을 작성하시오. 입력으로 주어진 행렬의 순서를 바꾸면 안 된다.

입력

첫째 줄에 행렬의 개수 N(1 ≤ N ≤ 500)이 주어진다.

둘째 줄부터 N개 줄에는 행렬의 크기 r과 c가 주어진다. (1 ≤ r, c ≤ 500)

항상 순서대로 곱셈을 할 수 있는 크기만 입력으로 주어진다.

출력

첫째 줄에 입력으로 주어진 행렬을 곱하는데 필요한 곱셈 연산의 최솟값을 출력한다. 정답은 2^31-1 보다 작거나 같은 자연수이다. 또한, 최악의 순서로 연산해도 연산 횟수가 2^31-1보다 작거나 같다.

풀이

행렬을 실제로 곱하는 순서를 전부 시도하면 경우의 수가 너무 많다. 대신 구간 [l, r]을 하나의 부분 문제로 보고, 그 구간을 어디서 나눌 때 비용이 최소가 되는지 구하는 구간 DP로 바꾸면 된다.

코드에서는 dp[i][j]를 i번째부터 j번째 행렬까지 곱하는 최소 연산 수로 두고, 중간 분할점 k를 기준으로 dp[i][k] + dp[k + 1][j] + arr[i][0] * arr[k + 1][0] * arr[j][1]를 계산한다. 마지막 곱셈 비용까지 포함해야 실제 연산 수가 완성된다.

짧은 구간부터 긴 구간 순으로 채우면, 큰 구간을 계산할 때 필요한 더 작은 구간 값이 이미 준비되어 있다. 결국 이 문제는 곱셈 자체보다 분할 위치를 어떻게 잡느냐가 핵심이다.

코드

java
import java.io.*;
import java.util.*;
 
public class Main {
 
    static int[][] arr;
 
    public static void main(String[] args) throws IOException {
        BufferedReader br = new BufferedReader(new InputStreamReader(System.in));
 
        int N = Integer.parseInt(br.readLine());
        arr = new int[N + 1][2];
 
        for (int i = 1; i <= N; i++) {
            StringTokenizer st = new StringTokenizer(br.readLine());
            arr[i][0] = Integer.parseInt(st.nextToken());
            arr[i][1] = Integer.parseInt(st.nextToken());
        }
 
        System.out.println(solve(N));
    }
 
    static int solve(int N) {
        int[][] dp = new int[N + 1][N + 1];
 
        // g: 범위, i: 시작, k: 중간, j: 끝
        for (int g = 1; g <= N; g++) {
            for (int i = 1; i + g <= N; i++) {
                int j = i + g;
                dp[i][j] = Integer.MAX_VALUE;
                
                for (int k = i; k < j; k++) {
                    dp[i][j] = Math.min(dp[i][j], dp[i][k] + dp[k + 1][j] + arr[i][0] * arr[k + 1][0] * arr[j][1]);
                }
            }
        }
 
        return dp[1][N];
    }
}

복잡도

  • 시간 복잡도: O(N3)O(N^3)
  • 공간 복잡도: O(N2)O(N^2)

마무리

행렬 곱셈 순서는 실제 곱셈보다도 어디서 끊을지를 정하는 문제가 핵심이다. 구간 DP로 바꾸면 복잡해 보이던 경우의 수가 점화식 하나로 정리된다.