몇 개의 요소로 이루어진 배열이 주어졌을 때, 배열을 회전시켜가며 얻을 수 있는 최대 가중치 합(maximum weighted sum)을 구하는 문제를 살펴보겠습니다. 배열 nums의 가중치 합은 다음과 같이 계산됩니다.
$$\mathrm{𝑆=\sum_{\substack{𝑖=1}}^{n}𝑖∗𝑛𝑢𝑚𝑠[𝑖]}$$
즉, 각 요소에 자신의 위치 인덱스(1부터 시작)를 곱한 값을 모두 더한 것이 가중치 합입니다.
문제 이해하기
예를 들어 입력이 L = [5, 3, 4]라면, 가능한 회전 상태별 가중치 합은 다음과 같습니다.
- 배열 [5, 3, 4]: 5 + 2×3 + 3×4 = 5 + 6 + 12 = 23
- 배열 [3, 4, 5]: 3 + 2×4 + 3×5 = 3 + 8 + 15 = 26 (최댓값)
- 배열 [4, 5, 3]: 4 + 2×5 + 3×3 = 4 + 10 + 9 = 23
따라서 정답은 26입니다.
효율적인 접근 방법
모든 회전 상태마다 처음부터 가중치 합을 다시 계산하면 O(n²)의 시간이 걸립니다. 하지만 회전이 일어날 때 가중치 합이 어떻게 변하는지 규칙을 파악하면 O(n) 만에 해결할 수 있습니다.
배열을 왼쪽으로 한 칸 회전하면 기존 요소들의 인덱스가 1씩 줄어들어 가중치 합에서 전체 요소의 합(sum_a)만큼 감소합니다. 동시에 맨 앞에 있던 요소는 맨 뒤로 이동해 인덱스 n의 가중치를 새로 얻게 되므로, nums[i] × n만큼 더해집니다. 이 점화식을 활용하는 것이 핵심입니다.
알고리즘 단계
- n := nums의 크기
- sum_a := nums의 모든 요소의 합
- ans := 초기 가중치 합(nums[i] × (i + 1)의 총합)
- cur_val := ans
- i를 0부터 n-1까지 반복:
- cur_val := cur_val − sum_a + nums[i] × n
- ans := ans와 cur_val 중 더 큰 값
- ans 반환
예제 코드
다음 구현을 통해 더 잘 이해해 보겠습니다.
def solve(nums): n = len(nums) sum_a = sum(nums) cur_val = ans = sum(nums[i] * (i + 1) for i in range(n)) for i in range(n): cur_val = cur_val - sum_a + nums[i] * n ans = max(ans, cur_val) return ans nums = [5,3,4] print(solve(nums))
입력
[5,3,4]
출력
26
이 알고리즘은 시간 복잡도 O(n), 공간 복잡도 O(1)로 동작하므로, 배열의 크기가 커져도 효율적으로 최대 가중치 합을 구할 수 있습니다.