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

Python으로 n개의 공 중 k개를 선택할 때 최댓값·최솟값 차이의 합 구하기

문제 이해하기

n개의 공이 있고, 각 공은 크기가 n인 배열 nums의 값으로 번호가 매겨져 있다고 가정해 보겠습니다. 즉, nums[i]는 i번째 공의 번호를 나타냅니다. 여기에 또 하나의 값 k가 주어지며, 매 차례마다 n개의 서로 다른 공 중에서 k개를 골라 그 번호들의 최댓값과 최솟값의 차이를 표에 기록합니다. 그다음 k개의 공을 다시 항아리에 넣고, 가능한 모든 조합을 선택할 때까지 이 과정을 반복합니다. 마지막으로 표에 기록된 모든 차이의 합을 구하되, 결과가 너무 커질 경우 109+7로 나눈 나머지를 반환하면 됩니다.

예시

예를 들어 입력이 n = 4, k = 3, nums = [5, 7, 9, 11]이라면 출력은 20이 됩니다. 가능한 조합은 다음과 같습니다.

  • [5, 7, 9] → 차이: 9 − 5 = 4
  • [5, 7, 11] → 차이: 11 − 5 = 6
  • [5, 9, 11] → 차이: 11 − 5 = 6
  • [7, 9, 11] → 차이: 11 − 7 = 4

따라서 4 + 6 + 6 + 4 = 20입니다.

접근 방법

모든 조합을 직접 만들어 차이를 더하는 방식은 조합의 수가 급격히 늘어나기 때문에 비효율적입니다. 대신 정렬된 배열에서 각 원소가 최댓값 또는 최솟값으로 등장하는 횟수를 조합론적으로 세면 선형 시간에 답을 구할 수 있습니다.

배열 nums가 오름차순으로 정렬되어 있다고 하면, 인덱스 i의 원소 nums[i]가 어떤 k개 조합에서 최댓값이 되려면 나머지 k−1개는 반드시 i보다 앞쪽 원소에서 골라야 하므로, 그런 조합의 수는 C(i, k−1)입니다. 같은 원리로 뒤쪽에서 세면 nums[n−1−i]가 최솟값이 되는 조합의 수 역시 C(i, k−1)입니다. 따라서 전체 합은 다음과 같이 정리됩니다.

Σ (nums[i] − nums[n−1−i]) × C(i, k−1)

조합 계수 C(i, k−1)는 매번 새로 계산하지 않고 점화식으로 업데이트하며, 나눗셈 연산은 모듈러 역원(inverse)을 미리 구해 처리합니다. 역원은 inv[i] = −(m // i) × inv[m % i] % m 공식을 이용해 선형 시간에 구할 수 있습니다.

알고리즘 단계

  • m := 109 + 7
  • inv := [0, 1]로 시작하는 모듈러 역원 리스트 생성
  • i를 2부터 n까지 순회하며 inv의 끝에 (m − (m // i) × inv[m mod i] mod m) 값 추가
  • comb_count := 1, res := 0으로 초기화
  • pick을 k−1부터 n−1까지 순회하며 다음을 수행:
    • res := res + (nums[pick] − nums[n−1−pick]) × comb_count mod m
    • res := res mod m
    • comb_count := comb_count × (pick + 1) mod m × inv[pick + 2 − k] mod m
  • res 반환

구현 코드

다음 구현을 통해 더 잘 이해해 보겠습니다.

def solve(n, k, nums):
    m = 10**9 + 7

    inv = [0, 1]
    for i in range(2, n + 1):
        inv.append(m - m // i * inv[m % i] % m)

    comb_count = 1
    res = 0
    for pick in range(k - 1, n):
        res += (nums[pick] - nums[n - 1 - pick]) * comb_count % m
        res %= m
        comb_count = comb_count * (pick + 1) % m * inv[pick + 2 - k] % m

    return res

n = 4
k = 3
nums = [5, 7, 9, 11]
print(solve(n, k, nums))

입력

4, 3, [5, 7, 9, 11]

출력

20

참고 사항

이 알고리즘은 nums가 오름차순으로 정렬되어 있다고 가정합니다. 입력 배열이 정렬되어 있지 않다면 sorted(nums)로 먼저 정렬한 뒤 함수에 전달해야 올바른 결과를 얻을 수 있습니다. 시간 복잡도는 O(n), 공간 복잡도는 역원 테이블 저장을 위해 O(n)입니다.