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)