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

파이썬으로 목표값에 가장 가까운 부분 수열의 합 찾기


배열 nums와 하나의 값 goal이 주어져 있다고 가정해 보겠습니다. 우리의 목표는 nums에서 부분 수열(subsequence)을 선택하여 그 합이 goal에 최대한 가깝도록 만드는 것입니다. 즉, 선택한 부분 수열의 합을 s라고 할 때, 절대 차이 |s - goal|를 최소화해야 합니다.

예를 들어, 입력이 nums = [8,-8,16,-1], goal = -3이라면 출력은 2가 됩니다. 부분 수열 [8,-8,-1]을 선택하면 합이 -1이 되고, 이때 절대 차이는 |-1 - (-3)| = 2로 가능한 최솟값입니다.

해결 접근 방법

모든 부분 수열을 완전 탐색하면 경우의 수가 지수적으로 늘어나기 때문에 비효율적입니다. 대신 절댓값이 큰 원소부터 처리하고, 남은 원소들의 합 범위를 활용해 유망하지 않은 탐색 경로를 미리 잘라내는(가지치기) 방식으로 문제를 해결할 수 있습니다. 구체적인 단계는 다음과 같습니다.

  • n := nums의 길이

  • nums를 각 요소의 절댓값을 기준으로 내림차순 정렬 (절댓값이 큰 원소를 먼저 처리해 조기 종료 효과를 극대화)

  • neg := 크기 n+1의 배열을 0으로 초기화

  • pos := 크기 n+1의 배열을 0으로 초기화

  • i를 n-1부터 0까지 1씩 감소시키며 반복:

    • nums[i] < 0이면:

      • neg[i] := neg[i+1] + nums[i]

      • pos[i] := pos[i+1]

    • 그렇지 않으면:

      • pos[i] := pos[i+1] + nums[i]

      • neg[i] := neg[i+1]

  • ans := |goal|

  • s := {0}을 담은 새 집합 생성

  • check(a, b) 함수 정의:

  • b < goal - ans 또는 goal + ans < a이면 False 반환, 그렇지 않으면 True 반환

메인 로직에서는 다음을 수행합니다.

  • i를 0부터 n-1까지 반복:

    • sl := s에서 check(x + neg[i], x + pos[i])가 참인 모든 x의 리스트

    • sl의 크기가 0이면 루프 탈출

    • s := sl로부터 새 집합 생성

    • sl의 각 x에 대해:

      • y := x + nums[i]

      • |y - goal| < ans이면 ans := |y - goal|로 갱신

      • ans == 0이면 0을 즉시 반환 (더 좋은 답은 존재하지 않음)

      • y를 s에 삽입

  • ans 반환

여기서 check() 함수가 핵심 역할을 합니다. 인덱스 i 이후의 음수 누적합(neg)과 양수 누적합(pos)을 이용해, 현재 상태에서 어떤 선택을 하더라도 이미 찾은 최적 차이(ans)보다 나아질 수 없다면 해당 경로를 더 이상 탐색하지 않습니다. 이를 통해 탐색 공간이 크게 줄어들어 효율성이 향상됩니다.

구현 코드

아래 파이썬 구현 예제를 통해 더 자세히 이해해 보겠습니다.

def solve(nums, goal):
   n = len(nums)
   nums.sort(key=lambda x: -abs(x))
   neg = [0 for _ in range(n+1)]
   pos = [0 for _ in range(n+1)]
   for i in range(n-1, -1, -1):
       if nums[i] < 0:
           neg[i] = neg[i+1] + nums[i]
           pos[i] = pos[i+1]
       else:
           pos[i] = pos[i+1] + nums[i]
           neg[i] = neg[i+1]
   ans = abs(goal)
   s = set([0])

   def check(a, b):
       if b < goal - ans or goal + ans < a:
           return False
       return True

   for i in range(n):
       sl = [x for x in s if check(x+neg[i], x+pos[i])]
       if len(sl) == 0:
           break
       s = set(sl)
       for x in sl:
           y = x + nums[i]
           if abs(y - goal) < ans:
               ans = abs(y - goal)
           if ans == 0:
               return 0
           s.add(y)
   return ans

nums = [8,-8,16,-1]
goal = -3
print(solve(nums, goal))

입력

[8,-8,16,-1], -3

출력

2