@laigary.com~/interview/coding/15-3-sum.md$
$ cat ./coding/15-3-sum.md
[Coding]·2023-01-29·9 min read

15. 3 Sum

15. 3 Sum

給一個陣列,找出所有不重複的三個數字組合,其和等於 0。注意要的是「值的組合」而不是索引,而且同樣的組合不能出現兩次。

這篇接在 1. 2 Sum 之後 —— 那篇的「排序 + 對撞指針」解法就是這裡要用的積木。

思路

從 2 Sum 遞迴降階

於是這樣就可以開始擴展當 n > 2 的時候,該怎麼辦呢?

  • n == 1 的時候,沒有所謂的總和問題,所以答案為空集合。
  • n == 2 的時候,就是所謂的 2 Sum,可以透過對撞指針找到
  • n == 3 的時候,是不是可以把設定的目標先減去陣列中的第 i 個數字,得到一個新目標,接著在剩下的區間中,用新的目標找 2 Sum
  • n == k 的時候,就可以推導出,先將目標減去第 i 個數字得到一個新目標後,在接下來的區間中找出 k-1 Sum

所以我們就找到了基本情況,n == 1n == 2,並且可以透過遞迴的方式來找出 nSum。

這個降階的寫法比「三層迴圈」值錢的地方在於:同一份程式碼可以解 3 Sum、4 Sum、k Sum,只要換一個參數。面試被追問「那 4 Sum 呢」的時候,答案是「把 3 換成 4」。

為什麼一定要先排序

排序做了兩件事,兩件都不可或缺:

  1. 讓對撞指針成立 —— 有序才能靠「和太小就移左、太大就移右」單向收斂,理由見 167
  2. 讓去重變得簡單 —— 排序後相同的值一定相鄰,所以「跳過重複」只需要跟隔壁比,不需要另外開一個 set 存已經產生過的組合。這和 26. Remove Duplicates from Sorted Array 是完全同一個道理。

第 2 點是這題最容易寫錯的地方。去重要在兩個層級都做:

  • 外層:固定的那個數字如果和下一個相同,要跳過,否則會產生一模一樣的組合
  • 內層:對撞指針找到一組答案後,左右都要跳過所有相同值,否則同一組會被重複收集

解題方向

class Solution:
    def twoSum(self, nums: List[int], start: int, target: int) -> List[List[int]]:
        left = start
        right = len(nums) - 1
        res = []
        while left < right:
            total = nums[left] + nums[right]
            leftVal = nums[left]
            rightVal = nums[right]
            if total < target:
                left += 1
            elif total > target:
                right -= 1
            else: # s == target
                res.append([nums[left], nums[right]]) 
                while (left < right and leftVal == nums[left]):
                    left += 1
                while (left < right and rightVal == nums[right]):
                    right -= 1
        return res

    def nSum(self, nums, n, start, target):
        size = len(nums)
        res = []
        if n < 2 or size < n:
            return res
        if n == 2:
            return self.twoSum(nums, start, target)
        else:
            i = start
            while i < size:
                tuples = self.nSum(nums, n - 1, i + 1, target - nums[i])
                for tup in tuples:
                    tup.append(nums[i])
                    res.append(tup)
                while i < size - 1 and nums[i] == nums[i + 1]:
                    i += 1
                i += 1
        return res

    def threeSum(self, nums: List[int]) -> List[List[int]]:
        nums.sort()
        return self.nSum(nums, 3, 0, 0)

幾個要注意的點:

  • twoSum 多了一個 start 參數,因為遞迴下來時只能在「固定的那個數字之後」的區間找,不然會重複使用同一個元素。
  • twoSum 回傳的是值不是索引 —— 排序之後索引已經沒有意義了,而這題要的正是值。
  • nSum 裡的 size < n 提前返回:剩下的元素不夠湊出 n 個,直接放棄。
  • tup.append(nums[i]) 是把固定的那個數字補回到子問題的答案裡,所以組合的順序不是排序後的順序 —— 題目不在意順序,所以沒差。

同一份 nSum 換個入口就是 18. 4 Sum

    def fourSum(self, nums: List[int], target: int) -> List[List[int]]:
        nums.sort()
        return self.nSum(nums, 4, 0, target)

補充

為什麼不用 hash table? 1. 2 Sum 的正解是 hash table,但那是因為它要原陣列的索引所以不能排序。這題要的是不重複的值組合,排序反而是幫手 —— 它同時給了對撞指針和廉價的去重。題目要什麼決定了能不能排序,這是整個 n Sum 家族最重要的分水嶺。

整個家族1. 2 Sum(hash table)、167. Two Sum II(已排序,對撞指針)、15. 3 Sum18. 4 Sum(同一份 nSum)、454. 4 Sum II(四個獨立陣列,回到 hash table)。整套面試應對策略見 2 Sum 面試應對策略,模板見 Two Pointers 模板

複雜度

3 Sum

  • 時間 O(n2) — 外層固定一個數字 O(n),內層對撞指針 O(n);排序的 O(nlogn) 被這個量級吃掉了
  • 空間 O(1) — 不算輸出和排序本身,只用了幾個索引

k Sum(同一份 nSum

  • 時間 O(nk1) — 每多一層就多一個 O(n) 的外層迴圈,最內層永遠是 O(n) 的對撞指針
  • 空間 O(k) — 遞迴深度是 k,和陣列長度無關

其中 n 是陣列長度。所以 4 Sum 是 O(n3)、5 Sum 是 O(n4) —— 對撞指針省下的是最後那一層,把 O(nk) 的暴力解降到 O(nk1)