Chương 40 - Bitmask Dynamic Programming

Khi state là subset của tập nhỏ (≤ 20 phần tử), ta encode subset bằng bitmask int 32-bit và DP trên bitmask. State space O(2^n). Pattern cho các bài TSP-style, set cover, partition.

Mục tiêu chương

Sau chương này, bạn sẽ:

  • Bitmask state khi n ≤ 20-22 (vì 2^20 ≈ 10^6).
  • Iterate submask: sub = (sub - 1) & mask.
  • Pattern: TSP-style, set cover, partition.
  • Path reconstruction với parent table.

Khi nào dùng pattern này?

  • n ≤ 20-22 (vì 2^20 ≈ 10^6 còn fit).
  • Bài “partition / cover / assignment” với set.
  • Bài “visit all nodes” / Travelling Salesman.

Template code

from functools import cache

@cache
def dp(mask, *extra_state):
    if mask == 0:        # base case (hoặc full = (1 << n) - 1)
        return base_value
    return best([dp(mask ^ (1 << i), ...) for i in iterate_bits(mask)])

Bài tự luyện cuối chương

  • LC 691 - Stickers to Spell Word
  • LC 1494 - Parallel Courses II
  • LC 1799 - Maximize Score After N Operations

40.1 Partition to K Equal Sum Subsets (LC 698)

Đề bài

Cho numsk. Chia thành k subset có tổng bằng nhau?

Ví dụ

Input:  nums=[4,3,2,3,5,2,1], k=4
Output: True

Ràng buộc

  • 1 <= k <= len(nums) <= 16
  • 1 <= nums[i] <= 10^4

Clarifying questions

  • Sum không chia hết k? → Trả False.
  • k = 1? → Trả True (cả mảng = 1 subset).

Hướng tiếp cận

dp[mask] = boolean “subset bitmask có thể partition thoả”.

Total = sum(nums). Mỗi subset tổng = total / k.

Duyệt: cho mỗi mask, nếu dp[mask] true và current_subset_sum % target == 0, thử thêm các element chưa dùng.

Code Python 3

from functools import cache
from typing import List

class Solution:
    def canPartitionKSubsets(self, nums: List[int], k: int) -> bool:
        total = sum(nums)
        if total % k: return False
        target = total // k
        nums.sort(reverse=True)
        if nums[0] > target: return False
        n = len(nums)

        @cache
        def dfs(mask: int, current_sum: int) -> bool:
            if mask == (1 << n) - 1:
                return True
            for i in range(n):
                if mask & (1 << i): continue
                new_sum = current_sum + nums[i]
                if new_sum > target: continue
                next_sum = new_sum if new_sum < target else 0
                if dfs(mask | (1 << i), next_sum):
                    return True
            return False

        return dfs(0, 0)

Phân tích độ phức tạp

  • Thời gian: O(k · 2^n · n) worst.
  • Bộ nhớ: O(2^n) cho memo.

Bình luận

  • Trick next_sum = 0 khi full: “đóng” subset hiện tại, mở subset mới.

Bài tự luyện liên quan

  • LC 416 - Partition Equal Subset Sum (Chương 29.5)
  • LC 1723 - Find Minimum Time to Finish All Jobs

40.2 Shortest Path Visiting All Nodes (LC 847)

Đề bài

Input: graph: List[List[int]] - adjacency list của 1 graph vô hướng liên thông n đỉnh (đỉnh đánh số 0..n-1). Tìm path ngắn nhất (đếm cạnh) visit tất cả đỉnh; được xuất phát/kết thúc tại đỉnh bất kỳ, được lặp lại đỉnh.

Ví dụ

Input:  graph = [[1,2,3], [0], [0], [0]]
        (adjacency list undirected; graph[i] = các node kề node i; được xuất phát/kết thúc tại node bất kỳ)
Output: 4   (số cạnh ngắn nhất để visit tất cả node; được lặp lại node)

Ràng buộc

  • n == len(graph)
  • 1 <= n <= 12

Clarifying questions

  • n = 1? → Trả 0.

Hướng tiếp cận

BFS với state (node, visited_mask). Đáp án = số bước đầu tiên đạt được state (_, full_mask).

Code Python 3

from collections import deque
from typing import List

class Solution:
    def shortestPathLength(self, graph: List[List[int]]) -> int:
        n = len(graph)
        full = (1 << n) - 1
        if n == 1: return 0
        queue = deque((i, 1 << i, 0) for i in range(n))   # (node, mask, dist)
        visited = {(i, 1 << i) for i in range(n)}
        while queue:
            node, mask, dist = queue.popleft()
            if mask == full: return dist
            for nb in graph[node]:
                new_mask = mask | (1 << nb)
                if (nb, new_mask) not in visited:
                    visited.add((nb, new_mask))
                    queue.append((nb, new_mask, dist + 1))
        return -1

Phân tích độ phức tạp

  • Thời gian: O(2^n · n²) state × transition.
  • Bộ nhớ: O(2^n · n).

Bình luận

  • TSP-style trên unweighted graph. BFS đảm bảo min steps.

Bài tự luyện liên quan

  • LC 943 - Shortest Superstring (bài 40.4)
  • LC 1879 - Min XOR Sum of Two Arrays (bài 40.6)

40.3 Smallest Sufficient Team (LC 1125)

Đề bài

Cho req_skills và mảng people[i] = list of skill. Tìm subset người nhỏ nhất cover tất cả skill.

Ví dụ

Input:  req_skills=["java","nodejs","reactjs"], people=[["java"],["nodejs"],["nodejs","reactjs"]]
Output: [0,2]

Ràng buộc

  • 1 <= len(req_skills) <= 16
  • 1 <= len(people) <= 60

Clarifying questions

  • req_skills rỗng? → Trả [].
  • Một người cover hết? → Trả [người đó].

Hướng tiếp cận

DP bitmask trên skills. dp[mask] = smallest team cover skills trong mask.

dp[mask | skills_of_p] = min(dp[mask | skills_of_p], dp[mask] + [p]).

Code Python 3

from typing import List

class Solution:
    def smallestSufficientTeam(self, req_skills: List[str], people: List[List[str]]) -> List[int]:
        skill_idx = {s: i for i, s in enumerate(req_skills)}
        n = len(req_skills)
        full = (1 << n) - 1
        people_mask = []
        for p in people:
            mask = 0
            for s in p:
                if s in skill_idx:
                    mask |= 1 << skill_idx[s]
            people_mask.append(mask)

        dp: dict[int, list[int]] = {0: []}
        for i, mask in enumerate(people_mask):
            for cur_mask, team in list(dp.items()):
                new_mask = cur_mask | mask
                if new_mask == cur_mask: continue
                if new_mask not in dp or len(team) + 1 < len(dp[new_mask]):
                    dp[new_mask] = team + [i]
        return dp[full]

Phân tích độ phức tạp

  • Thời gian: O(2^k · m) với k = số skill, m = số người.
  • Bộ nhớ: O(2^k).

Bình luận

  • Bẫy: memo phải lưu cả mask và people đã pick.
  • Follow-up: LC 1434 (Hats) ngược lại - pick assignment.

Bài tự luyện liên quan

  • LC 691 - Stickers to Spell Word
  • LC 1434 - Hats (Chương 44.5)

40.4 Find the Shortest Superstring (LC 943)

Đề bài

Cho mảng strings. Tìm string ngắn nhất chứa tất cả các string trong mảng làm substring.

Ví dụ

Input:  words = ["alex", "loves", "leetcode"]
Output: "alexlovesleetcode"   (chuỗi ngắn nhất chứa mọi word làm substring;
                               có thể có nhiều đáp án, trả về 1 trong số đó)

Ràng buộc

  • 1 <= len(words) <= 12
  • 1 <= len(words[i]) <= 20

Clarifying questions

  • words = 1 string? → Trả words[0].
  • Có duplicate trong words? → Có thể; dedup trước.

Hướng tiếp cận

Bitmask DP kết hợp TSP. State (mask, last_string_idx) = string ngắn nhất visit set mask và kết thúc tại string last_idx.

Pre-compute overlap[i][j] = max overlap khi nối string i

  • j.

Code Python 3

from typing import List

class Solution:
    def shortestSuperstring(self, words: List[str]) -> str:
        n = len(words)
        # overlap[i][j] = max k sao cho words[i] kết thúc bằng prefix length-k của words[j].
        overlap = [[0] * n for _ in range(n)]
        for i in range(n):
            for j in range(n):
                if i != j:
                    for k in range(min(len(words[i]), len(words[j])), 0, -1):
                        if words[i].endswith(words[j][:k]):
                            overlap[i][j] = k
                            break

        INF = float('inf')
        # dp[mask][i] = length của superstring ngắn nhất visit set mask, kết thúc tại i.
        dp = [[INF] * n for _ in range(1 << n)]
        parent = [[-1] * n for _ in range(1 << n)]
        for i in range(n):
            dp[1 << i][i] = len(words[i])

        for mask in range(1 << n):
            for i in range(n):
                if not (mask & (1 << i)) or dp[mask][i] == INF:
                    continue
                for j in range(n):
                    if mask & (1 << j):
                        continue
                    new_mask = mask | (1 << j)
                    new_len = dp[mask][i] + len(words[j]) - overlap[i][j]
                    if new_len < dp[new_mask][j]:
                        dp[new_mask][j] = new_len
                        parent[new_mask][j] = i

        # Trace path: tìm end-state có length min.
        full = (1 << n) - 1
        last = min(range(n), key=lambda i: dp[full][i])

        # Reconstruct order ngược.
        order: list[int] = []
        mask = full
        cur = last
        while cur != -1:
            order.append(cur)
            prev = parent[mask][cur]
            mask ^= 1 << cur
            cur = prev
        order.reverse()

        # Build superstring.
        result = words[order[0]]
        for idx in range(1, len(order)):
            prev_i, cur_i = order[idx - 1], order[idx]
            result += words[cur_i][overlap[prev_i][cur_i]:]
        return result

Phân tích độ phức tạp

  • Thời gian: O(2^n · n²).
  • Bộ nhớ: O(2^n · n) cho dp + parent.

Bình luận

  • Bài Hard nhất chương - kết hợp Bitmask DP + overlap precompute + path tracing.
  • Time: O(2^n · n²), Space: O(2^n · n). Với n ≤ 12, fit.

Bài tự luyện liên quan

  • LC 847 - Shortest Path Visiting All Nodes (bài 40.2)
  • LC 691 - Stickers to Spell Word

40.5 Maximum Students Taking Exam (LC 1349)

Đề bài

Phòng thi m × n. seats[i][j] = . (ok) hoặc # (hỏng). 2 học sinh không được kề bên trái/phải hoặc 2 góc chéo (vì copy được). Max số học sinh ngồi.

Ví dụ

Input:  seats = [["#",".","#","#",".","#"],
                 [".","#","#","#","#","."],
                 ["#",".","#","#",".","#"]]
        ("." = ghế OK, "#" = ghế hỏng; HS xem được bài 4 hướng chéo & cạnh)
Output: 4   (số HS tối đa xếp được mà không ai gian lận)

Ràng buộc

  • seats.length == m, seats[0].length == n
  • 1 <= m <= 8, 1 <= n <= 8

Clarifying questions

  • Toàn ghế hỏng? → Trả 0.

Hướng tiếp cận

Row DP bitmask. dp[i][mask] = max ở hàng i với học sinh ngồi theo mask.

Mỗi mask phải valid (không 2 bit kề) và không đè ô hỏng.

Transition: dp[i][mask] = max(dp[i-1][prev]) + popcount(mask) với prev không xung đột chéo với mask.

Code Python 3

from typing import List

class Solution:
    def maxStudents(self, seats: List[List[str]]) -> int:
        m, n = len(seats), len(seats[0])
        # row_bad[i] = bitmask các ô '#' trong hàng i.
        row_bad = [0] * m
        for i in range(m):
            for j in range(n):
                if seats[i][j] == '#':
                    row_bad[i] |= 1 << j

        # Valid masks: không 2 bit kề nhau.
        valid = [m for m in range(1 << n) if (m & (m << 1)) == 0]

        dp = {0: 0}     # mask -> max students
        for i in range(m):
            new_dp = {}
            for mask in valid:
                if mask & row_bad[i]: continue
                cnt = bin(mask).count('1')
                best = 0
                for prev, prev_cnt in dp.items():
                    if (mask & (prev << 1)) or (mask & (prev >> 1)): continue
                    best = max(best, prev_cnt)
                new_dp[mask] = best + cnt
            dp = new_dp
        return max(dp.values(), default=0)

Phân tích độ phức tạp

  • Thời gian: O(m · 2^n · 2^n) worst.
  • Bộ nhớ: O(2^n).

Bình luận

  • Bẫy: check 2 mask không kề chéo: mask & (prev << 1)mask & (prev >> 1).
  • Follow-up: LC 1411 (Paint Grid) cùng row-by-row mask DP.

Bài tự luyện liên quan

  • LC 1349 - Maximum Students Taking Exam (bài này)
  • LC 1986 - Min Number of Work Sessions

40.6 Minimum XOR Sum of Two Arrays (LC 1879)

Đề bài

2 mảng độ dài n. Hoán vị nums2 để min Σ nums1[i] ^ nums2[π(i)].

Ví dụ

Input:  nums1=[1,2], nums2=[2,3]
Output: 2

Ràng buộc

  • 1 <= len(nums1) == len(nums2) <= 14
  • 0 <= nums[i] <= 10^7

Clarifying questions

  • n = 1? → Trả nums1[0] ^ nums2[0].

Hướng tiếp cận

Bitmask DP O(n · 2^n). dp[mask] = min sum khi đã ghép cặp các index trong mask (cùng số bit = số i đã xét) với cùng số element của nums2.

Code Python 3

from typing import List

class Solution:
    def minimumXORSum(self, nums1: List[int], nums2: List[int]) -> int:
        n = len(nums1)
        INF = float('inf')
        dp = [INF] * (1 << n)
        dp[0] = 0
        for mask in range(1 << n):
            i = bin(mask).count('1')
            if i >= n: continue
            for j in range(n):
                if mask & (1 << j): continue
                new_mask = mask | (1 << j)
                dp[new_mask] = min(dp[new_mask], dp[mask] + (nums1[i] ^ nums2[j]))
        return dp[(1 << n) - 1]

Phân tích độ phức tạp

  • Thời gian: O(2^n · n).
  • Bộ nhớ: O(2^n).

Bình luận

  • Tinh tế: i = popcount(mask) cho thấy element nums1 thứ i đang được match.

Bài tự luyện liên quan

  • LC 698 - Partition to K Equal Sum Subsets (bài 40.1)
  • LC 691 - Stickers to Spell Word

Tóm tắt chương & Quyết định

Bitmask cookbook

Operation Code
Bit i set? (mask >> i) & 1
Set bit i mask | (1 << i)
Clear bit i mask & ~(1 << i)
Toggle bit i mask ^ (1 << i)
Lowest set bit mask & -mask
Pop count bin(mask).count('1') hoặc int.bit_count() (Py 3.10+)
Iterate submasks sub = mask; while sub: ...; sub = (sub - 1) & mask (cuối cùng còn sub=0)
Iterate set bits while mask: i = (mask & -mask).bit_length() - 1; mask &= mask - 1

Constraint feasibility table

n Bitmask 2^n Common DP shape Time
≤ 20 ≤ 10⁶ dp[mask] (TSP, Assignment) O(2^n · n)
≤ 16 ≤ 65k dp[mask][i] (TSP với endpoint) O(2^n · n²)
≤ 22-25 ≤ 33M Cần tối ưu hoặc submask enum O(3^n) cho subset-sum DP

Shortest Superstring (LC 943) - duplicate words

  • Trước khi DP, loại bỏ từ là substring của từ khác. LC test cases không có duplicate hoàn toàn nhưng có “chuỗi chứa chuỗi” → phải lọc.

Maximum Students (LC 1349) - row mask conflict

Row mask  s  (bit i = 1 nếu có HS ở cột i):
- Valid trong row: s & (s << 1) == 0 (không 2 HS kế).
- Compatible với prev mask p: 
    (s & (p << 1)) == 0 AND (s & (p >> 1)) == 0
- Phải nằm trong cell allowed: s & broken_mask[row] == 0.