Skip to content

2321. Maximum Score Of Spliced Array

#236

Problem link

https://leetcode.com/problems/maximum-score-of-spliced-array/

Problem Summary

배열 2개가 주어지고, 배열의 일부 구간을 다른 배열과 swap 가능할 때 두 배열의 합 중 최대가 가장 큰 값을 구하는 문제.

Solution

딱 보니 DP라서 처음에는 top-down으로 풀었다.

DP[i][j] = i 인덱스에서 j 플래그로 swap했을 때 최대값.
(j==0: swap 안함, j==1: swap 하는중, j==2: swap 완료)

해당 최댓값은 배열 하나에만 계산되므로 배열 2개를 서로 바꿔서 한번 더 돌려주면 된다.

정답이 나오긴 하는데 2792 ms로 하위 5%의 성능..
결국 bottom-up 으로 다시 제출해서 성능 2배 개선, 메모리도 대폭 줄였다 (1차원 DP)

참고로 discuss를 보면 Kadane알고리즘이라는 것도 있는데... 어차피 DP도 O(n) 이라 성능 차이는 별로 없다.
(https://leetcode.com/problems/maximum-score-of-spliced-array/discuss/2199139/100-fasteror-Most-Detailed-solution-for-complete-dummiesor-O(n)-Kadane's-single-loop)

Source Code

Bottom-Up (1264 ms / 31.5MB)

from functools import lru_cache
from typing import List


class Solution:
    def maximumsSplicedArray(self, nums1: List[int], nums2: List[int]) -> int:
        def solve(list1, list2):
            dp = [0 for _ in range(3)]
            dp[0] = list1[0]
            dp[1] = list2[0]

            for i in range(1, len(list1)):
                dp[2] = max(dp[1] + list1[i], dp[2] + list1[i])
                dp[1] = max(dp[0] + list2[i], dp[1] + list2[i])
                dp[0] = dp[0] + list1[i]

            return max(dp)

        return max(solve(nums1, nums2), solve(nums2, nums1))

Top-Down (2792 ms / 517.1 MB)

from functools import lru_cache
from typing import List


class Solution:
    def maximumsSplicedArray(self, nums1: List[int], nums2: List[int]) -> int:

        @lru_cache(None)
        def solve(idx, flag, swap):
            if idx == len(nums1):
                return 0

            if flag == 0:
                return max(solve(idx + 1, 0, swap) + nums1[idx], solve(idx + 1, 1, swap) + nums2[idx])
            elif flag == 1:
                return max(solve(idx + 1, 1, swap) + nums2[idx], solve(idx + 1, 2, swap) + nums1[idx])
            else:
                return solve(idx + 1, 2, swap) + nums1[idx]

        ret0 = solve(0, 0, 0)
        nums1, nums2 = nums2, nums1
        ret1 = solve(0, 0, 1)
        return max(ret0, ret1)