[백준] 2042번: 구간 합 구하기 (JAVA)

인간몽쉘김통통·2025년 2월 22일

백준

목록 보기
84/92
post-thumbnail

문제

https://www.acmicpc.net/problem/2042

이해

최대 1,000,000개 N만큼의 수열이 주어진다. 수열에 대해서 2가지 쿼리가 필요하다.

  1. 숫자 변경 (1 <= M <= 10,000)
    특정 위치의 수열의 숫자를 변경한다.

  2. 구간 합 구하기 (1 <= K <= 10,000)
    수열에서 특정 범위의 구간 합을 구해야 한다.

접근

당연하게도 수열을 배열로 관리하면 안된다. 수열의 크기는 1,000,000이고 구간 합 쿼리를 실행할 때마다 계산을 위해 N만큼의 시간이 필요하다. 쿼리의 최대 횟수가 10,000이기 때문에 시간복잡도 초과한다.

그렇다면 누적합은 어떨까? 누적합은 구간합을 1의 시간으로 처리할 수 있다. 하지만 본 문제는 1번 업데이트 쿼리가 존재한다. 업데이트가 될 때마다 이후 누적합 배열을 갱신해야 하기 때문에 N만큼의 시간이 결국 필요하다. 1번 쿼리의 횟수가 최대 10,000이기 때문에 딱히 개선될 여지는 없다.

본 문제를 해결하기 위해서는 수열을 트리로 관리해야 한다. 세그먼트 트리가 수정이 잦는 데이터의 구간 합을 구할 때 굉장히 유리하기 때문에 해당 자료구조를 사용하면 된다.

+) 세그먼트 트리에 대해서는 다음 링크를 확인해주세요.
https://wikidocs.net/209446

세그먼트 트리는 완전이진트리를 활용하여 수열의 구간 합을 미리 가지고 있는 형태를 띈다. 전체 합을 가지는 1번 노드를 선두로 각 구간을 이등분하여 구간 합을 자식이 갖도록 한다. 따라서, 1번 노드의 첫번째 자식 2번 노드는 수열의 1번부터 N/2까지의 구간합을, 두번째 자식 3번 노드는 수열의 N/2 + 1부터 N까지의 구간합을 가지게 된다.


(출처: 위키독스 https://wikidocs.net/209446)

다음으로 쿼리에 대해서 생각해보자. 쿼리 1번 숫자 업데이트는 트리에서 리프노드를 변경하면 된다. 세그먼트 트리는 구간 합의 정보를 가지고 있기 때문에 기존의 숫자와 수정된 숫자의 차이만큼 모든 부모, 조상노드를 수정하면 된다.

가령, 윗 그림에서 수열의 1번 노드(트리에서는 16번)가 수정되면 그의 부모, 조상인 8, 4, 2, 1(트리 기준 번호)를 수정하면 된다.

쿼리 2번 구간합은 어떻게 구하면 될까? 여기서는 이분탐색이 필요하다. 특정 구간 (a, b)의 합을 구한다고 할 때, 세그먼트 트리를 탐색하여 현재 위치가 무슨 범위의 구간 합을 의미하는지 판단한다.

구간끼리를 비교하기 때문에 여러 경우가 발생한다. 편의상 현재 노드의 범위를 (left, right)라고 하자.

  1. 현재 범위가 타겟 범위 밖일 때 (right < a) || (left > b)
  2. 완전히 포함될 때 (a <= left && right <= b)
  3. 그 외 (부분 포함될 때)

1번의 경우는 구하려는 대상 범위가 아니기 때문에 필요없다.
2번의 경우에는 현재 위치를 그대로 정답에 포함하면 된다.
3번의 경우에는 현재 위치의 자식 노드를 탐색해야 한다. 자식 노드를 탐색한다는 뜻은 현재 범위를 2등분하여 다시 판단한다는 뜻이다. (left, right)를 반으로 쪼개 다시 (a, b)와의 비교를 통해 구간 합에 포함할 지 계산하는 것이다.

이분탐색을 통해 구간 합을 최종적으로 구할 수 있기 때문에 시간복잡도를 logN으로 처리할 수 있다. 이것이 세그먼트를 사용하는 주된 이유이다.

풀이

1. 트리 생성

    private static void treeInit() {
        int idx = 1;
        while (Math.pow(2, idx) <= N) {
            idx++;
        }
        leafLayer = idx;
        treeSize = (int) Math.pow(2, leafLayer + 1);
        tree = new long[treeSize];
        leafStartIdx = (int) Math.pow(2, leafLayer);

        for (int i = 0; i < N; i++) {
            tree[leafStartIdx + i] = arr[i];
        }

        if (N > 1) {
            tree[1] = (DP(2) + DP(3));
        }
    }

초기화는 위와 같이 하였다. 필요한 완전이진트리를 생성하기 때문에 2의 제곱수의 크기로 배열을 생성했다. 리프 노드에 입력받은 수열이 들어가기 때문에 초기화 해주었다. 다음으로 세그먼트 트리의 부모들을 채우기 위해 탑다운 방식의 DP를 활용했다.

    private static long DP(int idx) {
        if (idx >= leafStartIdx) {
            return tree[idx];
        }

        if (tree[idx] == 0) {
            tree[idx] = DP(idx * 2) + DP(idx * 2 + 1);
        }

        return tree[idx];
    }

DP는 단순하게 작성했다. 현재 위치가 자식들을 단순히 더하면 된다.

2. update 쿼리

    private static void updateQuery(int a, long b) {
        int nodeIdx = (leafStartIdx + a - 1);
        long originNum = tree[nodeIdx];
        long diff = (b - originNum);

        tree[nodeIdx] = (long) b;
        int idx = (nodeIdx / 2);
        while (idx >= 1) {
            tree[idx] += diff;
            idx /= 2;
        }
    }

a의 위치가 리프 노드들 중에서 어디에 해당하는지 계산하고 원 숫자와의 차이를 계산한다. a의 부모, 조상들을 갱신하기 위해 반복문을 실행하여 diff만큼 더해주었다.

3. sum 쿼리

    private static Long sumQuery(int a, int b) {
        int right = (int) Math.pow(2, leafLayer);
        long ret = recursiveSum(1, 1, right, a, b);

        return ret;
    }

    private static Long recursiveSum(int n, int left, int right, int a, int b) {
        if (a <= left && right <= b) {
            return tree[n];
        } else if (right < a || left > b) {
            return (long) 0;
        } else {
            int mid = (left + right) / 2;
            return recursiveSum(n * 2, left, mid, a, b) + recursiveSum(n * 2 + 1, mid + 1, right, a, b);
        }
    }

재귀 함수는 구해야하는 타겟 범위와 현재 노드 숫자, 현재가 포함하는 범위를 파라미터로 전달하였다. 각 경우에 대한 분기처리는 위에서 설명한 것 처럼 처리했다.

부분적으로 포함될 때만 추가 이분 탐색이 필요하기 때문에 재귀처리하였다.

전체 코드

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

public class App {
    static int N, M, K;
    static long[] arr;
    static long[] tree;
    static int leafLayer;
    static int leafStartIdx;
    static StringBuilder sb = new StringBuilder();
    static int treeSize;

    public static void main(String[] args) throws Exception {
        BufferedReader br = new BufferedReader(new InputStreamReader(System.in));

        StringTokenizer st = new StringTokenizer(br.readLine());
        N = Integer.parseInt(st.nextToken());
        M = Integer.parseInt(st.nextToken());
        K = Integer.parseInt(st.nextToken());

        arr = new long[(int) N];
        for (int i = 0; i < N; i++) {
            arr[i] = Long.parseLong(br.readLine());
        }

        treeInit();

        // printForDebug();

        for (int i = 0; i < M + K; i++) {
            st = new StringTokenizer(br.readLine());

            int opt = Integer.parseInt(st.nextToken());
            int a = Integer.parseInt(st.nextToken());
            long b = Long.parseLong(st.nextToken());

            if (opt == 1) {
                updateQuery(a, b);
            } else {
                Long sum = sumQuery(a, (int) b);

                sb.append(sum).append("\n");
            }

            // printForDebug();
        }

        System.out.println(sb.toString());
    }

    private static void printForDebug() {
        System.out.println("\n========tree========");

        StringBuilder sbForDebug = new StringBuilder();
        for (int i = 0; Math.pow(2, i) < treeSize; i++) {
            for (int j = (int) Math.pow(2, i); j < Math.pow(2, i + 1); j++) {
                sbForDebug.append(tree[j]).append(" ");
            }
            sbForDebug.append("\n");
        }

        System.out.println(sbForDebug.toString());
    }

    private static void treeInit() {
        int idx = 1;
        while (Math.pow(2, idx) <= N) {
            idx++;
        }
        leafLayer = idx;
        treeSize = (int) Math.pow(2, leafLayer + 1);
        tree = new long[treeSize];
        leafStartIdx = (int) Math.pow(2, leafLayer);

        for (int i = 0; i < N; i++) {
            tree[leafStartIdx + i] = arr[i];
        }

        if (N > 1) {
            tree[1] = (DP(2) + DP(3));
        }
    }

    private static long DP(int idx) {
        if (idx >= leafStartIdx) {
            return tree[idx];
        }

        if (tree[idx] == 0) {
            tree[idx] = DP(idx * 2) + DP(idx * 2 + 1);
        }

        return tree[idx];
    }

    private static Long sumQuery(int a, int b) {
        int right = (int) Math.pow(2, leafLayer);
        long ret = recursiveSum(1, 1, right, a, b);

        return ret;
    }

    private static Long recursiveSum(int n, int left, int right, int a, int b) {
        if (a <= left && right <= b) {
            return tree[n];
        } else if (right < a || left > b) {
            return (long) 0;
        } else {
            int mid = (left + right) / 2;
            return recursiveSum(n * 2, left, mid, a, b) + recursiveSum(n * 2 + 1, mid + 1, right, a, b);
        }
    }

    private static void updateQuery(int a, long b) {
        int nodeIdx = (leafStartIdx + a - 1);
        long originNum = tree[nodeIdx];
        long diff = (b - originNum);

        tree[nodeIdx] = (long) b;
        int idx = (nodeIdx / 2);
        while (idx >= 1) {
            tree[idx] += diff;
            idx /= 2;
        }
    }
}

전체 코드는 위와 같다. 처음 푸는 세그먼트 트리 유형이라 내 스타일대로 작성했다. 전체적으로 불필요한 부분도 있기 때문에 최적화 한다면 코드를 많이 줄일 수 있을 것 같다.

결과

문제에서 다루는 숫자 범위가 Long형이어서 런타임에러가 발생했다.

profile
SW 0년차 개발자입니다.

0개의 댓글