posite 2026. 2. 22. 14:26

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

숫자가 변화하는 정수 배열의 구간합을 구하는 문제이다.

 

처음에는 단순하게 dp를 이용해서 구하려고 했으나, 배열의 크기와 쿼리의 수가 10^5 이기 때문에

시간 초과가 발생하여 제대로 된 풀이가 필요하다.

구간합을 구하고 변경하는데 O(log n) 의 시간만 소요되는 자료구조가 필요하며 이에 알맞는 자료구조는 세그먼트 트리이다.

세그먼트 트리는 각 구간의 누적합을 트리의 루트로 누적해가는 구조이며 예시는 다음과 같다.

루트(node=1): 전체 배열의 합
리프 노드: 개별 원소
내부 노드: 자식 노드들의 합
  
[1, 3, 5, 7, 9, 11]
                 [0~5] = 36
              /                \
        [0~2] = 9           [3~5] = 27
        /      \             /       \
   [0~1]=4  [2~2]=5    [3~4]=16   [5~5]=11
   /    \              /     \
[0]=1  [1]=3        [3]=7   [4]=9



이를 다음과 같은 규칙으로 제귀적으로 구축한다.

트리 노드 인덱스 규칙:
  - 루트:          node = 1
  - 왼쪽 자식:     2 * node
  - 오른쪽 자식:   2 * node + 1
  
public void build(long[] arr, int node, int start, int end) {
    if (start == end) {
        tree[node] = arr[start];
    } else {
        int mid = (start + end) / 2;
        build(arr, 2 * node, start, mid);
        build(arr, 2 * node + 1, mid + 1, end);
        tree[node] = tree[2 * node] + tree[2 * node + 1];
    }
}

 

구간의 합은 다음과 같이 구하였다.

node: 현재 노드 번호 start: 현재 시작점, end: 현제 종료점, l: 구하려는 구간 시작점, r : 구하려는 범위의 종료점

public long query(int node, int start, int end, int l, int r) {
    if (r < start || end < l) {
        return 0; // 범위 밖
    }
    if (l <= start && end <= r) {
        return tree[node]; // 완전히 포함
    }
    int mid = (start + end) / 2;
    long left = query(2 * node, start, mid, l, r);
    long right = query(2 * node + 1, mid + 1, end, l, r);
    return left + right;
}

 

변경은 해당 노드의 상위 노드의 모든 값을 수정해야 한다.

public void update(int node, int start, int end, int index, long val) {
    if (start == end) {
        tree[node] = val;
    } else {
        int mid = (start + end) / 2;
        if (index <= mid) {
            update(2 * node, start, mid, index, val);
        } else {
            update(2 * node + 1, mid + 1, end, index, val);
        }
        tree[node] = tree[2 * node] + tree[2 * node + 1];
    }
}

 

최종 코드는 다음과 같다.

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

public class 커피숍1275 {
    
    static long[] board;
    
    public static void main(String[] args) throws Exception {
        BufferedReader br = new BufferedReader(new InputStreamReader(System.in));
        StringTokenizer st = new StringTokenizer(br.readLine());
        int n = Integer.parseInt(st.nextToken()), q = Integer.parseInt(st.nextToken());
        board = new long[n + 1];
        st = new StringTokenizer(br.readLine());
        for (int i = 0; i < n; i++) {
            board[i] = Long.parseLong(st.nextToken());
        }
        
        SegmentTree tree = new SegmentTree(n);
        tree.build(board, 1, 0, n - 1);
        StringBuilder sb = new StringBuilder();
        for (int i = 0; i < q; i++) {
            st = new StringTokenizer(br.readLine());
            int x = Integer.parseInt(st.nextToken()), y = Integer.parseInt(st.nextToken());
            int start = Math.min(x, y), end = Math.max(x, y);
            int a = Integer.parseInt(st.nextToken());
            long b = Long.parseLong(st.nextToken());
            sb.append(tree.query(1, 0, n - 1, start - 1, end - 1)).append("\n");
            tree.update(1, 0, n - 1, a - 1, b);
        }
        br.close();
        System.out.print(sb);
    }
    
    static class SegmentTree {
        
        private long[] tree;
        
        public SegmentTree(int n) {
            this.tree = new long[4 * n];
        }
        
        // 초기 배열로 트리 구성
        public void build(long[] arr, int node, int start, int end) {
            if (start == end) {
                tree[node] = arr[start];
            } else {
                int mid = (start + end) / 2;
                build(arr, 2 * node, start, mid);
                build(arr, 2 * node + 1, mid + 1, end);
                tree[node] = tree[2 * node] + tree[2 * node + 1];
            }
        }
        
        // 구간 합 쿼리: [l, r]
        public long query(int node, int start, int end, int l, int r) {
            if (r < start || end < l) {
                return 0; // 범위 밖
            }
            if (l <= start && end <= r) {
                return tree[node]; // 완전히 포함
            }
            int mid = (start + end) / 2;
            long left = query(2 * node, start, mid, l, r);
            long right = query(2 * node + 1, mid + 1, end, l, r);
            return left + right;
        }
        
        // 점 업데이트: index 위치를 val로 변경
        public void update(int node, int start, int end, int index, long val) {
            if (start == end) {
                tree[node] = val;
            } else {
                int mid = (start + end) / 2;
                if (index <= mid) {
                    update(2 * node, start, mid, index, val);
                } else {
                    update(2 * node + 1, mid + 1, end, index, val);
                }
                tree[node] = tree[2 * node] + tree[2 * node + 1];
            }
        }
    }
}