416. Partition Equal Subset Sum
416. Partition Equal Subset Sum
給一個只含正整數的陣列,問能不能把它分成兩個子集合,讓兩個子集合的總和相等。
這篇混了我兩次寫這題的筆記,老實說我去看我之前寫的筆記是寫得很好,但是我必須承認我後來沒有想到一些細項。
思路
首先我們要確立目標,目標是給定一個陣列能不能分成兩個子集合,讓這兩個子集合的總和相等。
第一個要去想的地方比較簡單,既然兩個子集合要相等,那一個子集合一定是原陣列總和的一半,我們的目標就是要找到一個子集合,其總和是原陣列總和的一半。
另外,有可能原陣列的元素總和是奇數:例如:[1, 2, 1, 1] 這樣的陣列總和是 5 ,一定不會有辦法分成兩個集合,所以可以先用這樣的方式去剪枝(把所有總和為奇數的直接排除掉)。
接著是要怎麼選擇,每一個元素我們都可以有「選」或是「不選」這兩種方式,如果選的話,我們的目標就會減少當下的數值,如果不選的話,目標並不會改變,一直到如果我們選了最後一個數字,如果最後目標剛好等於 0 ,那就是我們透過選擇,找到了目標。
我再次寫這個題目的時候,其實我對於 Dynamic Programming 的熟悉程度已經好很多了,尤其我喜歡用自頂向下的方式去思考,這個題目也可以換一個角度想成,如果我走遍每個位置的數字,要考慮的是我要把這個數字放到 subset 1 or subset 2 。
這兩個角度其實是同一件事,差別只在要記什麼。
這樣的關係可以用一個遞迴的方式來呈現。
解題方向
記錄兩個子集合的總和
所以我寫出了:
class Solution:
def canPartition(self, nums: List[int]) -> bool:
@cache
def dp(i, sub1, sub2):
if i == len(nums):
return sub1 == sub2
one = dp(i + 1, sub1 + nums[i], sub2)
two = dp(i + 1, sub1, sub2 + nums[i])
return one or two
return dp(0, 0, 0)
不過這樣會有記憶體超過的問題,所以面試的時候可能會被問 follow up ,所以我們要去想怎麼優化。
這時候我才開始想起來,如果今天的總和是奇數,那不管怎麼分,都不可能有結果,所以要補上:
if sum(nums) % 2 == 1:
return False
我認為在面試的時候,如果有提到這一點,已經算是滿好的地步了,因為核心邏輯已經正確,但是如果用 Leetcode 去跑,這樣還是不夠,這時候就是要看單純程式碼的優化了,不過這在工作上的確滿常出現的,那就是 Short-circuit evaluation,例如:我們如果已知一個條件如果滿足了,後面的所有操作都可以不用執行,常見的情況類似 validation。
if condition == False:
return
# skip rest of execution
在這個題目,為了方便閱讀,我使用了 one and two 分別作為集合 1 and 2 ,這樣的情況就是兩個集合的總和我一定要遞迴走完才可以,但是我們在遞迴的回傳值,只要其中一個滿足了,就可以回傳 True ,也就是說,我們能夠推論出,當我 one 那條路如果可以回傳 True ,其實我就能直接回傳了,可以這樣寫:
class Solution:
def canPartition(self, nums: List[int]) -> bool:
if sum(nums) % 2 == 1:
return False
@cache
def dp(i, sub1, sub2):
if i == len(nums):
return sub1 == sub2
one = dp(i + 1, sub1 + nums[i], sub2)
if one:
return True
two = dp(i + 1, sub1, sub2 + nums[i])
return two
return dp(0, 0, 0)
或是我可以直接不要變數,透過程式 Eval 的順序,直接這樣寫:
class Solution:
def canPartition(self, nums: List[int]) -> bool:
if sum(nums) % 2 == 1:
return False
@cache
def dp(i, sub1, sub2):
if i == len(nums):
return sub1 == sub2
return dp(i + 1, sub1 + nums[i], sub2) or dp(i + 1, sub1, sub2 + nums[i])
return dp(0, 0, 0)
到這邊為止,這兩個優化蓋的是不同的情況:奇偶剪枝砍掉的是總和為奇數的輸入,短路砍掉的是「答案已經是 True、卻還在算另一條路」。如果答案是 False,兩條路本來就都得走完才能確定,短路是幫不上忙的。
只記錄一個子集合的總和
我在寫筆記的過程中,的確又想到了一些優化,那就是我們都需要算出總和,既然知道總和了,就可以除以二之後知道一個 subset 的總和應該是多少。
這時候我們其實並不需要知道 subset 1 和 subset 2 的總和各是多少,問題就可以變成:我當前的數字,要不要選進去?
class Solution:
def canPartition(self, nums: List[int]) -> bool:
total = sum(nums)
if total % 2 == 1:
return False
target = total // 2
@cache
def dp(i, sub):
if i == len(nums):
return sub == target
pick = dp(i + 1, sub + nums[i])
if pick:
return True
not_pick = dp(i + 1, sub)
return not_pick
return dp(0, 0)
遞迴
class Solution:
def canPartition(self, nums: List[int]) -> bool:
total = sum(nums)
if total % 2 != 0:
return False
def helper(i, target):
if target < 0:
return False
if i == len(nums):
return target == 0
else:
return helper(i + 1, target - nums[i]) or helper(i + 1, target)
return helper(0, total//2)
既然遞迴的方式已經寫好了,上面這個函式很明顯的存在著重疊的子問題,例如最後一個位置的「選」或「不選」,在前面幾個元素在做選擇時,會一直重複的計算,因此可以用記憶法,來記憶著已經做過的選擇。而每個位置的選擇或不選擇,可以透過當前位置與目標來決定。
自頂向下
class Solution:
def canPartition(self, nums: List[int]) -> bool:
total = sum(nums)
if total % 2 != 0:
return False
@cache
def helper(i, target):
if target < 0:
return False
if i == len(nums):
return target == 0
return helper(i + 1, target - nums[i]) or helper(i + 1, target)
return helper(0, total//2)
自底向上
一直以來我都覺得自底向上比較難想到,這題也是不太好想。
原陣列的每個位置都會有選擇或是不選擇兩種,自底向上的想法是,如果我們知道第 i 個位置在選擇的時候,是不是可以達到某個比我們的目標還要小的數值。
文字敘述很抽象,換成一個數學的例子就是,現在最後的目標是 5 ,前一個位置如果已經可以達到目標 5 了,那我這個位置就不用選擇,如果我的當前值是 2 ,那在前一個位置如果可以達到 5 - 2 = 3 ,那我就要選擇當前位置的元素,這一串選擇就是 dp[i][j] = dp[i - 1][j] or dp[i - 1][j - nums[i-1]] 這行在做的事情。
一開始的 base case 則是如果說一開始的目標就是 0 ,那不管有沒有選都是符合條件了。
class Solution:
def canPartition(self, nums: List[int]) -> bool:
total = sum(nums)
if total % 2 != 0:
return False
M = len(nums)
N = total // 2
dp = [[False] * (N + 1) for _ in range(M+1)]
for i in range(M+1):
dp[i][0] = True
for i in range(1, M+1):
for j in range(1, N+1):
if j - nums[i - 1] < 0:
dp[i][j] = dp[i - 1][j]
else:
dp[i][j] = dp[i - 1][j] or dp[i - 1][j - nums[i-1]]
return dp[M][N]
複雜度
以下用 表示陣列長度、 表示所有元素的總和。
遞迴
- 時間 — 每個元素都有選與不選兩條路,沒有記憶就會整棵樹展開
- 空間 — 遞迴深度
自頂向下(記兩個子集合的總和)
- 時間 — 狀態是(位置, 已選總和),每個狀態 ;第三個參數不會多出狀態,因為兩個總和相加一定等於前
i個元素的和 - 空間 — cache 的大小,另外還有 的遞迴堆疊
自頂向下(只記一個總和)
- 時間 — 狀態數跟上面那版一樣,少的是每個 key 的大小
- 空間 — 同上
自底向上
- 時間 — 填滿
(M+1) x (N+1)的表格,每格 - 空間 — 表格大小,沒有遞迴堆疊
這裡的 是數值大小而不是輸入長度,所以是 pseudo-polynomial,跟複雜度速查裡背包那一列是同一件事。