ALGORITHM NOTE1

BOJ 7453 - 합이 0인 네 정수

네 개 고르면 시간 터질텐데...

#algorithm#boj#gold#sorting#binary-search#two-pointers#meet-in-the-middle
아카이브로 돌아가기

문제 링크

문제

정수로 이루어진 크기가 같은 배열 A, B, C, D가 있다.

A[a], B[b], C[c], D[d]의 합이 0인 (a, b, c, d) 쌍의 개수를 구하는 프로그램을 작성하시오.

입력

첫째 줄에 배열의 크기 n (1 ≤ n ≤ 4000)이 주어진다. 다음 n개 줄에는 A, B, C, D에 포함되는 정수가 공백으로 구분되어져서 주어진다. 배열에 들어있는 정수의 절댓값은 최대 2^28이다.

출력

합이 0이 되는 쌍의 개수를 출력한다.

풀이

네 수를 한꺼번에 고르면 너무 크기 때문에, 두 수씩 합쳐 보는 중간에서 만나기 기법이 핵심이다. A+B의 모든 합과 C+D의 모든 합을 구해 두면, 두 합의 합이 0이 되는 경우만 세면 된다.

한쪽 합 배열을 정렬한 뒤 다른 쪽 합을 보며 -(A + B)가 몇 번 나오는지 이분 탐색으로 세거나, 두 배열을 정렬해 투 포인터처럼 셀 수 있다. 이렇게 하면 O(N4)O(N^4)O(N2logN)O(N^2 \log N) 수준으로 줄일 수 있다.

네 수 선택을 두 덩어리 합 비교로 바꾸는 것이 핵심이다.

코드

java
import java.io.*;
import java.util.*;
 
public class Main {
 
    static int N;
    static int[] A, B, C, D;
 
    public static void main(String[] args) throws IOException {
        BufferedReader br = new BufferedReader(new InputStreamReader(System.in));
        StringTokenizer st = new StringTokenizer(br.readLine());
 
        N = Integer.parseInt(st.nextToken());
        A = new int[N]; B = new int[N];
        C = new int[N]; D = new int[N];
 
        for (int i = 0; i < N; i++) {
            st = new StringTokenizer(br.readLine());
            A[i] = Integer.parseInt(st.nextToken());
            B[i] = Integer.parseInt(st.nextToken());
            C[i] = Integer.parseInt(st.nextToken());
            D[i] = Integer.parseInt(st.nextToken());
        }
 
        System.out.println(solve());
    }
 
    static long solve() {
        int size = N * N;
        int[] sumA = new int[size];
        int[] sumB = new int[size];
 
        int idx = 0;
        for (int i = 0; i < N; i++) {
            for (int j = 0; j < N; j++) {
                sumA[idx] = A[i] + B[j];
                sumB[idx] = C[i] + D[j];
                idx++;
            }
        }
 
        Arrays.sort(sumA);
        Arrays.sort(sumB);
 
        long ans = 0;
        int left = 0;
        int right = size - 1;
 
        while (left < size && right >= 0) {
            int curA = sumA[left];
            int curB = sumB[right];
            long sum = (long) curA + curB;
 
            if (sum == 0) {
                long cntA = 0, cntB = 0;
 
                while (left < size && curA == sumA[left]) {
                    cntA++;
                    left++;
                }
 
                while (right >= 0 && curB == sumB[right]) {
                    cntB++;
                    right--;
                }
 
                ans += cntA * cntB;
            }
            else if (sum < 0)
                left++;
            else
                right--;
        }
 
        return ans;
    }
}

복잡도

  • 시간 복잡도: O(N2logN)O(N^2 \log N)
  • 공간 복잡도: O(N2)O(N^2)

마무리

네 수를 직접 고르지 말고 둘씩 묶어 보면 문제가 훨씬 작아진다. 중간에서 만나기가 가장 잘 드러나는 문제 중 하나다.