[BOJ/JAVA] P1197 최소 스패닝 트리

아연·2023년 6월 22일

Algorithm

목록 보기
7/12

알고리즘 수업 시간에 배운 프림 알고리즘을 토대로 도전 ( •̀ ω •́ )✧!!
하였지만 메모리 초과로 실패해버림ㅎㅎ;;


문제 설명

그래프가 주어졌을 때, 그 그래프의 최소 스패닝 트리를 구하는 프로그램을 작성하시오.

최소 스패닝 트리는, 주어진 그래프의 모든 정점들을 연결하는 부분 그래프 중에서 그 가중치의 합이 최소인 트리를 말한다.


INPUT & OUTPUT

INPUT

  • 첫째 줄에 정점의 개수 V(1 ≤ V ≤ 10,000)와 간선의 개수 E(1 ≤ E ≤ 100,000)가 주어진다.
  • 다음 E개의 줄에는 각 간선에 대한 정보를 나타내는 세 정수 A, B, C가 주어진다. 이는 A번 정점과 B번 정점이 가중치 C인 간선으로 연결되어 있다는 의미이다.
    • C는 음수일 수도 있으며, 절댓값이 1,000,000을 넘지 않는다.
  • 그래프의 정점은 1번부터 V번까지 번호가 매겨져 있고, 임의의 두 정점 사이에 경로가 있다. 최소 스패닝 트리의 가중치가 -2,147,483,648보다 크거나 같고, 2,147,483,647보다 작거나 같은 데이터만 입력으로 주어진다.

예제 입력 1

3 3
1 2 1
2 3 2
1 3 3

OUTPUT

  • 첫째 줄에 최소 스패닝 트리의 가중치를 출력한다.

예제 출력 1

3

STRATEGY

📍Prim's Algorithm

프림 알고리즘으로 최소 스패닝 트리 문제를 해결하기 위해서는

  • nearest : (현재까지) 가장 가까운 노드를 저장한 배열
  • distance :가장 가까운 노드와의 거리를 저장한 배열

이 필요하다.

⚠️caution

distance[idx] == idxnearest[idx] 사이의 거리
이해가 잘 안된다면 참고하길 ⇒ 내 마음에 쏙 든 시각적 자료

⇒ 그래서 SOLUTION 부분의 코드를 보면

sum += W[vnear][nearest[vnear]]; 
//vnear노드 ~ vnear과 가장 가까운 노드 사이의 거리

distance[vnear] = -1;
//vnear노드는 이제 방문 안하도록

for (int j = 2; j <= v; j++) {
	if (W[j][vnear] < distance[j]) {
	distance[j] = W[j][vnear];
	nearest[j] = vnear;
    }
    //vear노드에서 갈 수 있는 최단거리로 distance & nearest 배열 업데이트
}

요 부분이 의미하는 게 앞서 말한 부분이다 !!


SOLUTION

import java.io.BufferedReader;
import java.io.IOException;
import java.io.InputStreamReader;
import java.util.StringTokenizer;

public class Main {

    public static void main(String[] args) throws IOException {
        BufferedReader br = new BufferedReader(new InputStreamReader(System.in));
        StringTokenizer st = new StringTokenizer(br.readLine());
        int v = Integer.parseInt(st.nextToken());
        int e = Integer.parseInt(st.nextToken());

        int[][] W = new int[v + 1][v + 1];
        int[] nearest = new int[v + 1];
        int[] distance = new int[v + 1];

        for (int i = 1; i <= v; i++) {
            st = new StringTokenizer(br.readLine());
            int e1 = Integer.parseInt(st.nextToken());
            int e2 = Integer.parseInt(st.nextToken());
            int d = Integer.parseInt(st.nextToken());
            W[e1][e2] = d;
            W[e2][e1] = d;
        }

        for (int i = 2; i <= v; i++) {
            nearest[i] = 1;
            distance[i] = W[1][i];
        }

        int vnear = 1;
        int sum = 0;
        for (int i = 1; i < v; i++) {
            int min = Integer.MAX_VALUE;
            for (int j = 2; j <= v; j++) {
                if (distance[j] >= 0 && distance[j] < min) {
                    min = distance[i];
                    vnear = j;
                }
            }
            sum += W[vnear][nearest[vnear]];
            distance[vnear] = -1;
            for (int j = 2; j <= v; j++) {
                if (W[j][vnear] < distance[j]) {
                    distance[j] = W[j][vnear];
                    nearest[j] = vnear;
                }
            }
        }
        System.out.println(sum);
    }
}

SOLUTION_with Priority Queue

import java.io.*;
import java.util.*;

public class Main {
    static int total;
    static List<Node>[] list;
    static boolean[] visited;
    static class Node implements Comparable<Node>{
        int vertex;
        int value;

        public Node(int vertex, int value) {
            this.vertex = vertex;
            this.value = value;
        }

        @Override
        public int compareTo(Node o) {
            return this.value - o.value;
        }
    }

    public static void main(String[] args) throws IOException{
        BufferedReader br = new BufferedReader(new InputStreamReader(System.in));
        StringTokenizer st = new StringTokenizer(br.readLine());
        int v = Integer.parseInt(st.nextToken());
        int e = Integer.parseInt(st.nextToken());

        list = new ArrayList[v + 1];
        visited = new boolean[v + 1];
        for(int i = 1; i < v + 1; i++) {
            list[i] = new ArrayList<>();
        }

        for(int i = 0; i < e; i++) {
            st = new StringTokenizer(br.readLine());
            int from = Integer.parseInt(st.nextToken());
            int to = Integer.parseInt(st.nextToken());
            int w = Integer.parseInt(st.nextToken());
            list[from].add(new Node(to, w));
            list[to].add(new Node(from, w));
        }

        prim(1);
        System.out.println(total);
    }

    static void prim(int start) {
        Queue<Node> pq = new PriorityQueue<>();

        pq.add(new Node(start, 0));
        while(!pq.isEmpty()) {
            Node p = pq.poll();
            int node = p.vertex;
            int weight = p.value;

            if(visited[node]) continue;
            visited[node] = true;
            total += weight;

            for(Node next : list[node]) {
                if(!visited[next.vertex]) {
                    pq.add(next);
                }
            }
        }
    }
}

0개의 댓글