Skip to content

416. Partition Equal Subset Sum

#181

Problem link

https://leetcode.com/problems/partition-equal-subset-sum/

Problem Summary

합이 같아지게 배열을 둘로 나눌 수 있는지 판단하는 문제.

Solution

처음엔 투포인터로 하다가 2, 2, 1, 1같은 케이스를 보고 dp로 풀었다.
핵심 아이디어는 합이 같아지게 나눈다는 것은 합이 배열의 전체 합의 절반이 되도록 만들 수 있는가를 판단하는 문제가 된다는 것이다.
이 아이디어로부터 배열의 각 원소를 더해보면서 합이 전체 합 / 2가 되는지 보면 된다.

통과했지만 틀린 코드였다

나중에 다시 보니 처음 제출은 Accepted인데 오답이었다... [4, 9, 1, 4, 3, 4, 5]에서 False가 나오는데 9 + 1 + 5 = 15라서 답은 True다. (더 짧은 반례는 [3, 4, 5, 10, 1, 5])

원인은 메모를 current_sum 하나로만 저장한 것.
solve(idx, s)의 답은 idx 뒤에 남은 원소에 따라 달라지는데, cached[current_sum]으로만 저장하면 깊은 idx에서 실패한 합이 원소가 더 남은 얕은 idx에서 그대로 재사용된다.

반례가 드물어서 테스트 케이스에 안 걸린 듯? 랜덤 4000개로는 하나도 안 나오고 6만 번 정도 찾아야 겨우 하나 나온다. Accepted라고 정답은 아니다.

다시 짜기

그래서 질문 방향을 바꿨다.

solve(idx, s)   s에서 시작해서 idx 뒤 원소로 target을 만들 수 있나   (남은 원소에 의존)
dp              앞 i개로 만들 수 있는 합들                           (이미 쌓인 사실)

앞을 쌓는 DP는 남은 원소에 의존하지 않아서 1차원으로 줄일 수 있다. 즉, 차원 수가 아니라 무엇을 저장하느냐가 문제였다.

원소는 한 번씩만 써야 해서 순회 중에 방금 추가한 합을 다시 쓰면 안 된다.
처음에는 층마다 copy.deepcopy로 복사했는데 안에 든 게 int랑 True밖에 없어서 완전 낭비였다... 새로 만든 합을 리스트에 모았다가 나중에 합쳐주면 복사 자체가 필요 없다.

n = 200, 값 50~100 기준으로 0.47초에서 0.09초로 줄었다. 시간복잡도는 O(n * target)

참고로 정수 비트로 하면 시프트 한 줄로 가능한 합을 한꺼번에 밀 수 있어서 훨씬 빠르다 (같은 입력에서 0.0001초).

bits = 1
for x in nums:
    bits |= bits << x
return (bits >> target) & 1 == 1

Source Code

다시 짠 것

from typing import List


class Solution:
    def canPartition(self, nums: List[int]) -> bool:
        if sum(nums) % 2 != 0:
            return False

        n = len(nums)
        target = sum(nums) // 2

        dp = set([0])

        for i in range(n):
            newSum = []
            for s in dp:
                if s + nums[i] > target:
                    continue

                newSum.append(s + nums[i])

            for s in newSum:
                dp.add(s)

        return target in dp

처음 제출 (Accepted, 오답)

from typing import List


class Solution:
    def canPartition(self, nums: List[int]) -> bool:
        s = sum(nums)

        if s % 2 != 0:
            return False

        target = s // 2
        dp = [False] * (target + 1)
        cached = [False] * (s + 1)

        def solve(idx, current_sum):
            if current_sum == target:
                return True
            if current_sum > target:
                return False
            if idx == len(nums):
                return False
            if cached[current_sum]:
                return dp[current_sum]

            dp[current_sum] = solve(idx + 1, current_sum + nums[idx]) or solve(idx + 1, current_sum)
            cached[current_sum] = True
            return dp[current_sum]

        return solve(0, 0)