문제 이해하기
배열 nums에 n개의 숫자가 주어져 있다고 가정해 봅시다. 우리는 배열에서 두 숫자로 이루어진 쌍(pair)을 선택해야 하며, 다음 조건을 만족해야 합니다.
두 숫자의 배열 내 위치 차이 = 두 숫자의 합
n개의 원소를 가진 배열에서 만들 수 있는 전체 쌍의 개수는 n(n - 1)/2개입니다. 이 중에서 위 조건을 만족하는 쌍이 총 몇 개인지 구하는 것이 목표입니다.
예를 들어 입력이 n = 8, nums = {4, 2, 1, 0, 1, 2, 3, 3}이라면 출력은 13이 됩니다. 즉, 이 배열에는 조건을 만족하는 쌍이 총 13개 존재합니다.
접근 방법
이 문제는 다음 단계를 따라 해결할 수 있습니다.
길이가 n인 배열 vals를 정의한다.
i := 0부터 i < n일 때까지 반복하면서(i는 1씩 증가):
vals[i] := i + 1 - nums[i]
배열 vals를 오름차순으로 정렬한다.
res := 0으로 초기화한다.
i := 0부터 i < n일 때까지 반복하면서(i는 1씩 증가):
k := nums[i] + i + 1
res := res + (vals 배열에서 k보다 큰 값이 처음 나타나는 위치 - k 이상인 값이 처음 나타나는 위치)
res를 반환한다.
핵심 아이디어
각 원소의 위치를 1부터 시작하는 인덱스로 생각하면 조건식을 깔끔하게 변형할 수 있습니다. 인덱스 j가 i보다 클 때 위치 차이는 (j + 1) − (i + 1)이므로 다음과 같이 정리됩니다.
(j + 1) − (i + 1) = nums[i] + nums[j]
→ j + 1 − nums[j] = i + 1 + nums[i]
왼쪽 변을 vals[j], 오른쪽 변을 k로 정의하면, 조건을 만족하는 쌍을 찾는 문제는 "정렬된 vals 배열에서 k와 같은 값을 가지는 원소의 개수 세기"로 바뀝니다. vals를 미리 정렬해 두면 lower_bound와 upper_bound(이분 탐색)를 활용해 각 k에 대해 일치하는 원소의 개수를 O(log n)에 구할 수 있으므로, 전체 시간 복잡도는 O(n log n)이 됩니다.
C++ 구현 예제
아래 구현을 통해 더 자세히 이해해 보겠습니다.
#include <bits/stdc++.h>
using namespace std;
int solve(int n, vector<int> nums){
vector<int> vals(n);
for(int i = 0; i < n; i++)
vals[i] = i + 1 - nums[i];
sort(vals.begin(), vals.end());
int res = 0;
for(int i = 0; i < n; i++) {
int k = nums[i] + i + 1;
res += upper_bound(vals.begin(), vals.end(), k) - lower_bound(vals.begin(), vals.end(), k);
}
return res;
}
int main() {
int n = 8;
vector<int> nums = {4, 2, 1, 0, 1, 2, 3, 3};
cout << solve(n, nums);
return 0;
}
입력
8, {4, 2, 1, 0, 1, 2, 3, 3}
출력
13