动态规划状态压缩
状态压缩动态规划:位运算表示集合状态、旅行商问题(TSP)、棋盘覆盖与排列型 DP。
1. 状态压缩原理
1.1 核心思想
当状态中包含集合信息时,用二进制位表示集合元素的存在与否:
集合 {0, 2, 4} → 二进制 10101 → 十进制 21
位 i = 1 → 元素 i 在集合中
位 i = 0 → 元素 i 不在集合中
n 个元素的子集数: 2^n
1.2 位运算操作
# 添加元素 i
S | (1 << i)
# 删除元素 i
S & ~(1 << i)
# 检查元素 i 是否在集合中
(S >> i) & 1
# 集合大小(1的个数)
bin(S).count('1')
# 枚举所有子集
for sub = S; sub > 0; sub = (sub - 1) & S
# 最低位的1
S & (-S)
2. 旅行商问题(TSP)
2.1 问题描述
给定 个城市和两两之间的距离,求从城市 0 出发经过所有城市恰好一次并返回的最短路径。
2.2 状态定义
2.3 状态转移
2.4 实现
def tsp(dist):
n = len(dist)
full = (1 << n) - 1
INF = float('inf')
dp = [[INF] * n for _ in range(1 << n)]
dp[1][0] = 0 # 从城市0出发,只经过城市0
for S in range(1, 1 << n):
for i in range(n):
if not (S & (1 << i)):
continue
for j in range(n):
if S & (1 << j):
prev = S ^ (1 << i)
if dp[prev][j] < INF:
dp[S][i] = min(dp[S][i], dp[prev][j] + dist[j][i])
# 回到起点
ans = INF
for i in range(1, n):
ans = min(ans, dp[full][i] + dist[i][0])
return ans
2.5 复杂度
状态数: 2^n × n
转移: O(n) 每个状态
总时间: O(2^n × n^2)
空间: O(2^n × n)
n ≤ 20 时可行(2^20 ≈ 100万)
3. 棋盘覆盖问题
3.1 轮廓线 DP
用二进制表示当前行的轮廓线状态:
def domino_tiling(m, n):
# m行n列棋盘,1×2骨牌覆盖
full = (1 << n) - 1
dp = [0] * (1 << n)
dp[full] = 1 # 初始:上一行全部填满
for row in range(m):
for col in range(n):
new_dp = [0] * (1 << n)
for S in range(1 << n):
if dp[S] == 0:
continue
# 当前格已被上方骨牌覆盖
if S & (1 << col):
new_dp[S ^ (1 << col)] += dp[S]
# 水平放置骨牌
if col + 1 < n and not (S & (1 << col)) and not (S & (1 << (col + 1))):
new_dp[S | (1 << (col + 1))] += dp[S]
dp = new_dp
return dp[full]
4. 排列型 DP
4.1 全排列问题
def permutation_dp(n, condition):
"""计算满足条件的排列数"""
dp = [0] * (1 << n)
dp[0] = 1
for S in range(1 << n):
pos = bin(S).count('1') # 已放置的位置数
for i in range(n):
if not (S & (1 << i)) and condition(pos, i):
dp[S | (1 << i)] += dp[S]
return dp[(1 << n) - 1]
5. 优化技巧
5.1 滚动数组
# 只保留当前层和上一层
dp_prev = [0] * (1 << n)
dp_curr = [0] * (1 << n)
# 交替使用,空间从 O(2^n × n) 降至 O(2^n)
5.2 枚举优化
# 枚举 S 的子集
for sub = S; sub > 0; sub = (sub - 1) & S:
# 处理 sub
pass
# 枚举 S 的超集
sup = S
while sup < (1 << n):
# 处理 sup
sup = (sup + 1) | S
5.3 预处理合法状态
# 预先过滤不合法状态,减少枚举
valid_states = [S for S in range(1 << n) if is_valid(S)]