BOJ_행렬곱셈순서_11049 (Java)

융바오·2025년 1월 19일

Problem Solving

목록 보기
44/89

문제 링크

성능 요약

메모리: 18568 KB, 시간: 236 ms

분류

다이나믹 프로그래밍

제출 일자

2025년 1월 18일 21:06:02

문제 설명

크기가 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)

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

출력

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

풀이

느낀점

  • 행렬문제만 나오면 어지러운데, 다시 규칙을 찾아보니 행렬을 잘 몰라도 문제에서 규칙을 이해하면 되는 문제였다.
  • 처음에는 그리디로 쉽게 풀었다가 0%에서 틀려버렸다. 그리디로 풀기에는 반례가 너무 많았다.
  • 알고리즘 분류가 다이나밍프로그래밍이라는 걸 보고 고민해봐도 방향성을 잡기 힘들었다 (DP최약체 ㅠ)
  • 풀이설명을 흐린눈으로 참고해서 1차로 풀고 또 실패해서, 2차에는 아예 코드를 참고해서 디버깅 했다.
  • DP에서 모든 경우의 수를 고려할 수 있는 방법인지 검토해보아야 한다는 걸 깨달았다.

설계 : 40분

  • (AxB)*(BxC) 행렬의 곱 결과 행렬의 크기는 (AXC)이다.
  • (AxB)*(BxC)*(CxD) 전체 행렬 곱 결과 크기는 (AxD)이다. 즉, 곱한 행렬 범위에서 가장 앞 행렬의 행과 가장 뒤 행렬의 열만 남는다.
  • (AxB)*(BxC)*(CxD)(DxW)*(WxY)*(YxZ) 두 행렬곱의 결과를 서로 곱하는 연산 수는 A*D*Z 이다. 즉, 각 행렬곱 범위에서 가장 앞 행렬의 행과 가장 뒤 행렬의 열만 가지고 연산 수를 구할 수 있다.
  • 가장 짧은 범위부터 각 범위의 행렬 곱에 대해 최소값을 기록하며 dp테이블을 채워보기로 한다.
  • dp[n][n]테이블을 2차원 배열로 start 행렬부터 end 행렬까지의 곱의 최소 연산 횟수를 기록한다.
  • 가장 작은 범위는 start부터 start+1 행렬까지의 곱이고, 가장 넓은 범위는 0부터 n-1행렬까지의 곱이다.
  • 유의할 점: 3개 이상의 행렬을 곱할때, 해당 범위를 두개로 나누는 위치에 따라 연산 수가 달라진다.
    • 즉, 범위를 나누는 인덱스를 모두 순회하며 최소값으로 저장해야한다.
    • 작은 범위만 생각해보고 범위를 늘릴때 앞뒤로 하나씩 추가하는 모양만 고려해서 틀렸었다.

코드(Java)

  • 구현 시간: 100분
/**
 * Author: yngbao97, Yuk Yejin
 * Problem: 행렬 곱셈 순서_11049
 * Date: 2025.01.18
 */

import java.lang.reflect.Array;
import java.util.*;
import java.lang.*;
import java.io.*;

public class Main {
	static BufferedReader br;
	static BufferedWriter bw;
	static StringTokenizer st;

	public static void main(String[] args) throws Exception {

		br = new BufferedReader(new InputStreamReader(System.in));
		bw = new BufferedWriter(new OutputStreamWriter(System.out));

        int n = Integer.parseInt(br.readLine());
        int[][] matrix = new int[n][2];
        int[][] dp = new int[n][n];

        for (int i = 0; i < n; i++) {
            st = new StringTokenizer(br.readLine(), " ");
            matrix[i][0] = Integer.parseInt(st.nextToken());
            matrix[i][1] = Integer.parseInt(st.nextToken());
        }

        for (int l = 1; l < n; l++) {
            for (int start = 0; start < n-l; start++) {
                int end = start + l;
                dp[start][end] = Integer.MAX_VALUE;
                for (int mid = start; mid < end; mid++) {
                    dp[start][end] = Math.min(dp[start][end],
                                        dp[start][mid] + dp[mid+1][end] + matrix[start][0] * matrix[mid][1] * matrix[end][1]);
                }
            }
        }

        bw.write(String.valueOf(dp[0][n-1]));
		bw.flush();
		bw.close();
		br.close();
	}
}

0개의 댓글