https://www.acmicpc.net/problem/2261
공부 날짜 : 2024.12.22
정답 참조 여부 : O
2차원 평면상에 n개의 점이 주어졌을 때 이 점들 중 가장 가까운 두 점을 구하는 프로그램을 작성하시오
뭐 언제나 그렇지만 가장 간단한 방법은 모든 좌표의 점을 비교해서 최소값을 찾는 방법
이다 하지만 플레 2 문제에서 그런 단순한 문제를 줄 리 없으니 당연히 n의 최대가 10만이였고, 다른 방법을 찾아야 한다.
해당 문제가 1차원 선이였으면 아마 정렬을 한 뒤에 좌우만 비교했으면 됐을 것이다. 거기에 착안해서 먼저 x축을 기준으로 정렬을 해줬는데 이후 방법이 애매 했다.
결론적으로는 비교하는 기준을 분할정복으로 구간을 나누는 것인데 이럴 경우 3가지 경우로 나뉜다.
그래서 해당 문제를 풀기 위해서 구간을 나눠서 1번과 2번을 구해준다.
그 다음 중앙선을 기준으로 x축 거리가 1번,2번에서 구한 거리보다 짧은 점들에 대해서만 따로 모아 최단 거리를 비교(가지치기)함으로 써 시간을 줄 일 수 있었다.
이번 문제에서는 특별히 특정 객체(Node)를 다른 기준으로(x축, y축)정렬 할 필요성이 있었다.
기존에 내가 구했던 병합정렬은 객체 내의 operator<(){}에 대해서만 정렬 할 수 있었으므로 이를 위해 비교함수를 넘겨받아 정렬하되, 비교함수를 제시하지 않으면 객체 내의 operator<(){}함수로 비교하도록 업그래이드 했다.
#if 1
#define _CRT_SECURE_NO_WARNINGS
#define LLINF 8999999999999999999
#include <iostream>
int abs_(int a) {
return a < 0 ? -a : a;
}
struct Node {
int x;
int y;
bool operator<(const Node& other) const {
return x < other.x;
}
}node[100000];
template<typename T, typename Compare>
void merge_sort(int left, int right, T *arr, Compare comp) {
if (left >= right) return;
int mid = (left + right) >> 1;
merge_sort(left, mid, arr, comp);
merge_sort(mid + 1, right, arr, comp);
int i = left;
int j = mid + 1;
int size = right - left + 1;
T* temp = new T[size];
int k = 0;
while (i <= mid && j <= right) {
if (comp(arr[i], arr[j])) temp[k++] = arr[i++];
else temp[k++] = arr[j++];
}
while (i <= mid) temp[k++] = arr[i++];
while (j <= right) temp[k++] = arr[j++];
for (int t = 0; t < size; ++t) {
arr[left + t] = temp[t];
}
delete[] temp;
}
// Compare 인자를 받지 않는 오버로드
// 여기서 기본적으로 a < b를 비교하는 람다를 사용하여 기본 정렬 기준 제공
template<typename T>
void merge_sort(int left, int right, T *arr) {
merge_sort(left, right, arr, [](const T &a, const T &b) {
return a < b;
});
}
long long dist(const Node &a, const Node &b) {
return ((long long)a.x - b.x) * (a.x - b.x) + ((long long)a.y - b.y) * (a.y - b.y);
}
long long closestdist(int left, int right, Node* arr) {
// 쌍비교에서 3개이면 3회비교로 최소 비교 만족
if (right - left + 1 <= 3) {
long long min_value = LLINF;
for (int i = left; i <= right; ++i)
for (int j = i + 1; j <= right; ++j) {
long long d = dist(arr[i], arr[j]);
min_value = min_value > d ? d : min_value;
}
return min_value;
}
int mid = (left + right) >> 1;
long long a = closestdist(left, mid, arr);
long long b = closestdist(mid + 1, right, arr);
long long min_value = a > b ? b : a;
Node *temp = new Node[right - left + 1];
int size = 0;
for (int i = left; i <= right; ++i) {
if ((node[i].x - node[mid].x) * (node[i].x - node[mid].x) < min_value)
temp[size++] = node[i];
}
merge_sort(0, size - 1, temp, [](const Node& a, const Node& b) {
return a.y < b.y;
});
for (int i = 0; i < size; ++i) {
for (int j = i + 1; j < size && temp[j].y - temp[i].y < min_value; ++j) {
long long min_temp = dist(temp[i], temp[j]);
min_value = min_value > min_temp ? min_temp : min_value;
}
}
delete[] temp;
return min_value;
}
long long solve() {
int N;
std::cin >> N;
for (int i = 0; i < N; ++i)
std::cin >> node[i].x >> node[i].y;
merge_sort(0, N - 1, node);
return closestdist(0, N - 1, node);
}
int main() {
std::ios_base::sync_with_stdio(0);
std::cin.tie(0); std::cout.tie(0);
std::cout << solve() << "\n";
}
#endif