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

Python으로 인접 요소 간 절대 차이가 k 이하인 가장 긴 부분 수열의 길이 구하기

숫자로 이루어진 리스트와 정수 k가 주어졌을 때, 모든 인접한 요소 간의 절대 차이가 k 이하인 가장 긴 부분 수열(subsequence)의 길이를 구하는 프로그램을 만들어 보겠습니다.

예를 들어 입력이 nums = [5, 6, 2, 1, -6, 0, -1], k = 4라면 출력은 6이 됩니다. 실제로 [5, 6, 2, 1, 0, -1]처럼 인접 요소 간 차이(|5−6|=1, |6−2|=4, |2−1|=1, |1−0|=1, |0−(−1)|=1)가 모두 4 이하인 길이 6짜리 부분 수열을 만들 수 있기 때문입니다.

접근 방법: 세그먼트 트리로 O(n log n) 최적화

단순 동적 계획법(DP)으로 풀면 O(n²)의 시간이 필요하지만, 세그먼트 트리(Segment Tree)좌표 압축, 이진 탐색(bisect)을 결합하면 O(n log n)으로 개선할 수 있습니다. 핵심 아이디어는 다음과 같습니다.

  • 배열의 값을 정렬해 각 값의 순위(좌표)를 매기고, 세그먼트 트리에는 "해당 값으로 끝나는 가장 긴 부분 수열의 길이"를 저장합니다.
  • 각 요소 x를 처리할 때, [x−k, x+k] 범위에 속한 값들의 인덱스 구간을 bisect로 찾아 해당 구간의 최댓값을 조회합니다.
  • 조회한 최댓값에 1을 더해 현재 요소까지의 길이로 트리를 갱신(update)하고, 전체 정답을 업데이트합니다.

알고리즘 단계

이 문제를 해결하기 위해 다음 단계를 따릅니다.

  1. update(i, x) 함수 정의 — i := i + n으로 리프 노드 위치로 이동한 뒤, i가 0이 아닌 동안 segtree[i] := max(segtree[i], x)로 갱신하고 i := i // 2로 부모 노드로 올라갑니다.
  2. query(i, j) 함수 정의 — ans := −∞로 초기화하고, i := i + n, j := j + n + 1로 변환합니다. i < j인 동안 i가 홀수이면 ans := max(ans, segtree[i]) 후 i := i + 1을, j가 홀수이면 j := j − 1 후 ans := max(ans, segtree[j])를 수행합니다. 이후 i := i // 2, j := j // 2로 한 단계 위로 올라가고, 반복이 끝나면 ans를 반환합니다.
  3. 메인 로직 초기화 — nums = [5, 6, 2, 1, −6, 0, −1], k = 4로 설정하고, n := 2^(log₂(len(nums) + 1) + 1), segtree := [0] * 100000으로 선언합니다.
  4. snums := sorted(nums)로 정렬된 복사본을 만들고, index := {x: i for i, x in enumerate(snums)}로 값→순위 매핑 딕셔너리를 생성합니다.
  5. ans := 0으로 초기화한 뒤, nums의 각 요소 x에 대해 다음을 수행합니다.
    • lo := bisect_left(snums, x − k) — x−k 이상인 값이 시작되는 가장 왼쪽 삽입 위치
    • hi := bisect_right(snums, x + k) − 1 — x+k 이하인 값이 끝나는 가장 오른쪽 위치
    • count := query(lo, hi) — 해당 범위에서의 최대 부분 수열 길이 조회
    • update(index[x], count + 1) — 현재 요소를 포함한 길이로 트리 갱신
    • ans := max(ans, count + 1)
  6. 모든 요소를 처리한 후 ans를 반환합니다.

구현 예제

import math, bisect
class Solution:
   def solve(self, nums, k):
      n = 2 ** int(math.log2(len(nums) + 1) + 1)
      segtree = [0] * 100000
      def update(i, x):
         i += n
         while i:
            segtree[i] = max(segtree[i], x)
            i //= 2
      def query(i, j):
         ans = -float("inf")
         i += n
         j += n + 1
         while i < j:
            if i % 2 == 1:
               ans = max(ans, segtree[i])
               i += 1
            if j % 2 == 1:
               j -= 1
               ans = max(ans, segtree[j])
            i //= 2
            j //= 2
         return ans
      snums = sorted(nums)
      index = {x: i for i, x in enumerate(snums)}
      ans = 0
      for x in nums:
         lo = bisect.bisect_left(snums, x - k)
         hi = bisect.bisect_right(snums, x + k) - 1
         count = query(lo, hi)
         update(index[x], count + 1)
         ans = max(ans, count + 1)
      return ans
ob = Solution()
print(ob.solve([5, 6, 2, 1, -6, 0, -1], 4))

입력

[5, 6, 2, 1, -6, 0, -1], 4

출력

6

마무리

이 풀이는 세그먼트 트리의 범위 최댓값(range maximum query) 기능을 활용해, 각 요소마다 연결 가능한 이전 값들의 범위를 빠르게 탐색합니다. 좌표 압축 덕분에 값의 실제 크기가 아무리 커도 배열 크기는 요소 개수에 비례하므로 메모리도 효율적이며, 전체 시간 복잡도는 O(n log n)입니다. 값의 범위가 넓거나 입력 크기가 큰 코딩 테스트 문제에서 특히 유용하게 적용할 수 있는 패턴입니다.