Computer >> 컴퓨터 >  >> 프로그래밍 >> Python

파이썬으로 두 리스트에서 K개의 최대 합 쌍 찾기

두 개의 숫자 리스트 nums0nums1, 그리고 정수 k가 주어져 있다고 가정해 보겠습니다. 우리의 목표는 각 쌍이 nums0의 정수 하나와 nums1의 정수 하나로 구성되도록 하여, 합이 가장 큰 k개의 쌍을 찾는 것입니다. 그리고 선택된 모든 쌍의 합을 반환해야 합니다.

예를 들어 입력이 nums0 = [8, 6, 12], nums1 = [4, 6, 8], k = 2라고 한다면, 출력은 38이 됩니다. 가장 큰 쌍은 [12, 8]과 [12, 6]이며, 각각의 합은 20과 18이므로 전체 합은 38입니다.

문제 해결 접근 방식

이 문제는 최소 힙(min-heap)을 활용하면 효율적으로 해결할 수 있습니다. 두 리스트를 오름차순으로 정렬한 뒤, 각 리스트의 마지막 원소(가장 큰 값)끼리 만든 쌍부터 시작하여 힙에 넣습니다. 이후 힙에서 가장 큰 합을 가진 쌍을 하나씩 꺼내면서, 해당 쌍에서 인접한 후보 쌍(왼쪽 또는 아래 방향)을 힙에 추가하는 방식으로 상위 k개의 쌍을 순차적으로 탐색합니다.

구체적인 단계는 다음과 같습니다.

  • k > len(nums0) * len(nums1)인 경우, 만들 수 있는 쌍의 개수보다 요청된 개수가 많으므로 0을 반환합니다.
  • pq := 새로운 최소 힙을 생성합니다.
  • ans := 0으로 초기화합니다.
  • nums0과 nums1을 오름차순으로 정렬합니다.
  • i, j := 각각 len(nums0) − 1, len(nums1) − 1 (마지막 인덱스)로 설정합니다.
  • visited := 이미 힙에 넣은 좌표를 추적하기 위한 집합(set)을 생성합니다.
  • 힙 pq에 (−(nums0[i] + nums1[j]), i, j)를 푸시합니다. 최대 합을 구해야 하므로 부호를 반전시켜 저장합니다.
  • k번 반복하면서 다음을 수행합니다.
    • 힙 pq에서 (sum, i, j)를 팝합니다.
    • x := nums0[i − 1] + nums1[j]를 계산하고, 좌표 (i − 1, j)를 방문하지 않았다면 visited에 추가한 뒤 (−x, i − 1, j)를 힙에 푸시합니다.
    • y := nums0[i] + nums1[j − 1]을 계산하고, 좌표 (i, j − 1)를 방문하지 않았다면 visited에 추가한 뒤 (−y, i, j − 1)를 힙에 푸시합니다.
    • ans := ans + (−sum)으로 누적합니다.
  • ans를 반환합니다.

구현 코드

아래 파이썬 코드를 통해 동작 과정을 더 잘 이해할 수 있습니다.

from heapq import heappush, heappop

class Solution:
    def solve(self, nums0, nums1, k):
        if k > len(nums0) * len(nums1):
            return 0
        pq = []
        ans = 0
        nums0.sort(), nums1.sort()
        i, j = len(nums0) - 1, len(nums1) - 1
        visited = set()
        heappush(pq, (-(nums0[i] + nums1[j]), i, j))
        for _ in range(k):
            sum, i, j = heappop(pq)
            x = nums0[i - 1] + nums1[j]
            if not (i - 1, j) in visited:
                visited.add((i - 1, j))
                heappush(pq, (-x, i - 1, j))
            y = nums0[i] + nums1[j - 1]
            if not (i, j - 1) in visited:
                visited.add((i, j - 1))
                heappush(pq, (-y, i, j - 1))
            ans += -sum
        return ans

ob = Solution()
print(ob.solve([8, 6, 12], [4, 6, 8], 2))

입력

[8, 6, 12],[4, 6, 8],2

출력

38

핵심 포인트 정리

이 알고리즘의 시간 복잡도는 정렬에 O(n log n), 힙 연산에 O(k log k)가 소요되므로 전체적으로 O(n log n + k log k)입니다. 모든 가능한 쌍을 완전 탐색하는 O(n·m) 방식보다 k가 작을 때 훨씬 효율적이라는 점이 큰 장점입니다. 또한 visited 집합으로 중복 좌표를 걸러주지 않으면 같은 쌍이 여러 번 힙에 들어가 잘못된 결과가 나올 수 있으므로, 이 부분이 정확성의 핵심임을 기억해 두세요.