[자료구조] DisjointSet / 백준 1976 - 여행가자

메르센고수·2024년 2월 6일

DataStructure

목록 보기
2/2

개요

Union-Find 알고리즘을 사용하는 대표적인 자료구조가 DisjointSet이다.
그 중에서도 미로에 관해 다룰 때 DisjointSet을 사용하는데 '자료구조론' 수업 때 배운 내용으로 복습할 겸 리뷰를 해보겠다. 그 때 기억으로는 벽과 cell의 index를 구하는 과정이 상당히 힘들었던 것으로 기억한다.

참고

백준에도 DisjointSet 관련된 문제들이 꽤나 많다. 그 중에서 가장 대표적인 문제들 몇문제만 가져와 봤다.

  1. 백준 1717 - 집합의 표현 (Gold5)
  2. 백준 1976 - 여행 가자 (Gold4)
  3. 백준 1043 - 거짓말 (Gold4)
  4. 백준 20040 - 사이클 게임 (Gold4)
  5. 백준 1939 - 중량제한 (Gold3)

DisjointSet

DisjointSet (서로소 집합) - Wikipedia

주요한 연산으로는 Union과 Find 두 가지가 있다.
1. Union
: Union의 경우 두 집합의 대표 노드가 같은 집합인 경우는 그대로 두고, 다른 집합인 경우 두 집합을 하나로 병합하는 연산이다.

void union_(DisjointSets *S, int i, int j){
    int p1=find(S,i);
    int p2=find(S,j);

    if(p1!=p2){
        if(S->ptr_arr[p2]<S->ptr_arr[p1]){
            S->ptr_arr[p1]=p2;
        }else{
            if(S->ptr_arr[p2]==S->ptr_arr[p1])
                S->ptr_arr[p1]--;
            S->ptr_arr[p2]=p1;
        }
    }
}

대략 적인 틀은 위의 코드와 같은 틀이다. 해석을 해보면 먼저 find 함수를 통해 i와 j의 대표 노드를 찾은 다음 p1!=p2인 경우 (p1과 p2가 다른 집합에 있는 경우) 두 집합을 합치는데, 이 때 두 집합의 대표 노드의 크기를 비교해서 크기가 더 큰쪽으로 병합이 되도록 설정해준다.


  1. Find
    : Find 함수는 집합의 대표 노드를 찾는 함수이다.
int find(DisjointSets *S, int x){
    if(S->ptr_arr[x]<=0)
        return x;
    else    
        return find(S,S->ptr_arr[x]);
}

해석을 해보면, 집합의 vertex가 음수인 경우(대표 노드인 경우) x를 반환하고, 대표 노드를 찾지 못한 경우 다음 노드로 넘어가서 또 다시 find 연산을 수행하는 재귀함수의 형태로 구성되어 있다.

풀이


조금 난잡해 보이지만, 파란 글씨가 Cell의 번호이고 빨간 글씨가 벽의 번호이다.
n=6일 때의 예시이므로 n=6과 cell의 번호, 벽의 index를 적절히 조합해서 부숴야 하는 벽의 index를 찾고 제한 범위를 구해줘서 미로를 구성해 나가면 된다.

결론적으로 목표는 위의 예시 그림에서 처럼 8번 Cell과 14번 Cell을 같은 집합으로 만들어주기 위해 필요한 Cell과 벽의 index간의 관계를 찾아서 벽부수기를 진행하는 것이다. 그 과정을 1번 Cell과 마지막 36번 Cell이 같은 집합에 속하게 될 때까지 반복해서 미로의 통로를 열어주는 것이다.

소스 코드

#include <stdio.h>
#include <stdlib.h>
#include <time.h>

typedef struct _DisjointSet{
    int size;
    int *ptr_arr;
}DisjointSets;

typedef struct _PrintDisjointSet{
    int size;
    int *ptr_arr;
}PrintDisjointSets;

void init(DisjointSets* S, PrintDisjointSets* maze, int n);
int find(DisjointSets* S, int x);
void union_(DisjointSets* S, int i, int j);
void CreateMaze(DisjointSets* S, PrintDisjointSets* maze, int n);
void PrintMaze(DisjointSets*S, PrintDisjointSets* maze, int n);
void FreeMaze(DisjointSets* S, PrintDisjointSets* maze);

int main(int argc, char* argv[]){
    int num;
    FILE *fi = fopen(argv[1], "r");
    fscanf(fi, "%d", &num);
    fclose(fi);

    DisjointSets* S = (DisjointSets*)malloc(sizeof(DisjointSets));
    PrintDisjointSets* maze = (PrintDisjointSets*)malloc(sizeof(PrintDisjointSets));

    init(S, maze, num);
    CreateMaze(S, maze, num);
    PrintMaze(S, maze, num);
    FreeMaze(S, maze);
    return 0;
}


void init(DisjointSets *S, PrintDisjointSets *maze, int n){
    int total = n*n; // 정점의 개수
    S->size = total;
    S->ptr_arr = (int*)malloc(sizeof(int)*(total+1));

    maze->size = 2*n*(n+1); // 벽의 개수
    maze->ptr_arr = (int*)malloc(sizeof(int)*(maze->size));

    for(int i=0;i<=maze->size;i++){
        if(i==n||i==2*n*(n+1)-(n+1)){ // 출발점과 도착점은 0으로 초기화
            maze->ptr_arr[i] = 0;
        }else{
            maze->ptr_arr[i] = 1;
        }
    }

    for(int i=0;i<=total;i++){
        S->ptr_arr[i] = 0;
    }
}
int find(DisjointSets *S, int x){
    if(S->ptr_arr[x]<=0){
        return x;
    }else{
        return find(S, S->ptr_arr[x]);
    }
}
void union_(DisjointSets *S, int i, int j){
    int p1 = find(S,i);
    int p2 = find(S,j);

    if(p1 != p2){
        if(S->ptr_arr[p2] < S->ptr_arr[p1]){
            S->ptr_arr[p1] = p2;
        }else{
            if(S->ptr_arr[p2] == S->ptr_arr[p1])
                S->ptr_arr[p1]--;
            S->ptr_arr[p2] = p1;
        }
    }
}
void CreateMaze(DisjointSets *S, PrintDisjointSets *maze, int n){
    srand(time(NULL)); // random
    int i;
    int skip=0;
    int last_wall = n*(2*n+1)-1;

    while(find(S,1) != find(S,n*n)){ // 시작과 끝이 같은 집합이 될 때까지 반복
        skip = 0;
        int wall = rand()%(2*n*(n+1)); // 막혀있는 벽을 랜덤으로 구성

        for(i=0; i<n; i++){
            if(wall == i || wall == (2*n+1)*n+i || wall == n+(2*n+1)*i || wall == 2*n+(2*n+1)*i){ // 북, 남, 서, 동
                skip = 1;
                break;
            }
        }

        if(skip == 1)
            continue;
        if(wall % (2*n+1) < n){ // 가로벽
            int cell1 = (wall/(2*n+1)-1)*n+wall%(2*n+1)+1; // 위쪽 셀
            int cell2 = cell1 + n; // 아래쪽 셀
            printf("%d\n", cell1);
            if(find(S,cell1) != find(S,cell2)){
                union_(S, cell1, cell2);
                maze->ptr_arr[wall] = 0;
            }
        }else{ // 세로벽
            int cell1 = wall/(2*n+1)*n+wall%(2*n+1)-n; // 왼쪽 셀
            int cell2 = cell1 + 1; // 오른쪽 셀
            printf("%d\n", cell1);
            if(find(S,cell1) != find(S,cell2)){
                union_(S, cell1, cell2);
                maze->ptr_arr[wall] = 0;
            }
            if(wall == last_wall)
                break;
        }
    }
}
void PrintMaze(DisjointSets* S, PrintDisjointSets *maze, int n){
    for(int i=0; i<maze->size; i++){
        if(maze->ptr_arr[i] == 1){
            if(i%(2*n+1)<n)
                printf(" -");
            else
                printf("| ");
        }else{
            printf("  ");
        }

        if(i%(2*n+1)==n-1 || i%(2*n+1)==2*n) // 제일 오른쪽 벽 또는 제일 아래 벽
            printf("\n");
    }
}
void FreeMaze(DisjointSets *S, PrintDisjointSets *maze){
    free(S->ptr_arr);
    free(S);
    free(maze->ptr_arr);
    free(maze);
}

결과

27
25
3
34
34
21
7
25
32
35
14
21
23
22
15
14
15
28
20
10
18
2
33
14
29
32
5
1
29
19
22
24
26
17
9
23
1
28
12
4
17
11
5
11
1
32
14
28
26
13
2
1
19
22
30
13
2
 - - - - - -
      |     |
 -     -   -
|   |       |
 - - - - -
|       |   |
 -     - -  
|   | |   | |
 - -
|       |   |
   - -   - -
| |
 - - - - - -

미로 모양이 조금 허접하긴 하지만, 시작과 끝이 뚫려있기 때문에 제대로 구현이 되었다는 것을 알 수 있다.


백준 1976 - 여행가자

백준 1976 - 여행가자 (Gold4)
위에서 언급했던 백준 문제 중 Disjoint set(분리집합)을 이용해서 푸는 문제가 있어서 풀어보았다.

풀이

이 문제는 분리집합으로 풀어도 되고, DFS나 BFS 같은 그래프 탐색 알고리즘을 사용해서 풀어도 풀린다. C로 구현했던 DisjointSet의 기억을 떠올려서 구조체와 Union-Find 알고리즘을 활용하여 문제를 풀었다.

소스 코드

/*문제 : https://www.acmicpc.net/problem/1976
  알고리즘 : 자료구조, 그래프, disjoint set
  티어 : Gold4
*/

#include <iostream>
#include <vector>
#include <algorithm>
#define MAX 1001
using namespace std;

int N,M;

typedef struct DisjointSet{
    int size;
    int *arr;
}DisjointSet;

int Find(DisjointSet *set, int x){
    if(set->arr[x] == x){
        return x;
    }
    return Find(set, set->arr[x]);
}

void Union(DisjointSet *set, int x, int y){
    x = Find(set, x);
    y = Find(set, y);
    if (x < y)
        set->arr[y] = x;
    else
        set->arr[x] = y;
}

int main(void){
    ios::sync_with_stdio(false);
    cin.tie(NULL);
    cout.tie(NULL);

    cin >> N >> M;
    DisjointSet *set = new DisjointSet; // C의 malloc과 같은 기능
    set->size = (N+1)*(N+1);
    set->arr = new int[set->size];


    for(int i=1; i<=N; i++){
        set->arr[i] = i;
    }

	// 2차원 배열에 입력정보 저장
    for(int i=1; i<=N; i++){
        for(int j=1; j<=N; j++){
            cin >> set->arr[i*N+j];
        }
    }

    for(int i=1; i<=N; i++){
        for(int j=1; j<=N; j++){
            if(set->arr[i*N+j] == 1){
                Union(set, i, j);
            }
        }
    }
    
    int root;
    for (int i=0; i<M; i++){ 
        int x;
        cin>>x;
        if(i==0) // 시작점의 집합 찾기
            root=Find(set,x);
        else{
        	// 여행 계획으로 여행이 불가능한 경우
            if(Find(set, root) != Find(set, x)){
                cout << "NO";
                delete set;
                return 0;
            }
        }
    }
    cout << "YES";
    delete set;
    return 0;
}

결과

profile
블로그 이전했습니다 (https://phj6724.tistory.com/)

0개의 댓글