SOS-DP 详解
简介
SOS-DP,英文名 Sum over Subsets(SOS) dynamic programming,即子集和DP,用来求某一集合的所有子集对应状态之和。
给定一个含 $2^N$ 个元素的集合 $A$,下标用 $i$ 来表示,求一个 $F_{state}$,对应为:
$$
F_{state}=\sum_{i\subseteq state / i \text{&} state = state}A[i]
$$
解法:
Bruteforce
1
2
3
4
5
6
7for (int mask = 0; mask < (1 << N); mask++) {
for (int i = 0; i < (1 << N); i++) {
if ((mask & i) == i) {
F[mask] += A[i];
}
}
}对于每个 $mask$,暴力枚举哪些状态为它的子集,时间复杂度 $O(4^N)$
Suboptimal Solution
1
2
3
4
5
6
7
8// iterate over all the masks
for (int mask = 0; mask < (1 << N); mask++) {
F[mask] = A[0];
// iterate over all the subsets of the mask
for (int i = mask; i > 0; i = (i - 1) & mask) {
F[mask] += A[i];
}
}对于每个 $mask$,只会枚举它所包含的子集 $i$。
若该 $mask$ 含有 $k$ 个 $1$,则它所枚举的状态数量为 $2 ^ k$,时间复杂度为 $O(\sum_{k=0}^N \binom{N}{k}2^k=(1+2) ^ N = 3^N)$
SoS Dynamic Programming solution
1 | for (int i = 0; i < (1 << N); i++) { |
在 Suboptimal Solution 中,发现对于每个状态 $mask$,如果它包含有 $k$ 个 $0$,则会被枚举 $2^k$ 次(它对应 $2^k$ 个超集,每个超集都会枚举一次 $mask$ 状态)
考虑如何在不同的 $F_{state}$ 中建立联系:
设状态 $S_{state,i}={x|x\subseteq state \land state\oplus x < 2^{i+1}}$,即 $S_{state,i}$ 中包含的状态为仅有 0~i 位可以与 $state$ 状态不同,而其余位均相同的子集 $x$ 的集合
$F(S)$ 之间的状态转移为
$$
\begin{aligned}
& F_{state,0}=A[state]\
& F_{state,i}=
\begin{cases}
F_{state, i-1},&\text{if } state & (1 << i)=0\
F_{state,i-1}+F_{state\oplus (1<<i), i-1}, &\text{if }state & (1 << i)\ne 0
\end{cases}
\end{aligned}
$$
如果 $state$ 的第 $i$ 位为 $0$,则它只会包括第 $i$ 位为 $0$ 的子集;反之会包括第 $i$ 位为 $0$ 和 $1$ 的子集。
发现每层状态只会与上一层相关,经过空间压缩后得到上述代码。
该算法又被称为高维前缀和,其逆过程,也就是将F[mask] += F[mask ^ (1 << i)];改为F[mask] -= F[mask ^ (1 << i)];可得高维差分,一般先用高位前缀和,变换后再用高维差分,得到所需的严格等的信息。
例题:
-
在数组 $a$ 对于每一个元素 $a_i$,找出一个元素 $a_j$,满足 $a_i \text{&} a_j=0$,其中元素大小 $a_i\le 4\times 10^6$
将 $a_i$ 标记为状态,只需要看 $state=((1<<N)-1)\oplus a_i$ 对应状态的子集有无数组 $a$ 中元素即可,再用一位表示这个状态含有哪些子集,只需要一个,用一个维度添加特定子集即可
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46int main() {
std::ios::sync_with_stdio(false);
std::cin.tie(nullptr);
int n;
cin >> n;
std::vector<int> a(n);
for (int i = 0; i < n; i++) {
cin >> a[i];
}
constexpr int N = 22, M = (1 << 22);
std::vector<pii> dp(M);
for (int i = 0; i < M; i++) {
dp[i] = {0, i};
}
for (int i = 0; i < n; i++) {
dp[a[i]].fi = 1;
}
for (int i = 0; i < N; i++) {
for (int mask = 0; mask < M; mask++) {
if (mask >> i & 1) {
if (!dp[mask].fi) {
dp[mask].se = dp[mask ^ (1 << i)].se;
}
dp[mask].fi += dp[mask ^ (1 << i)].fi;
}
}
}
for (int i = 0; i < n; i++) {
int mask = (M - 1) ^ a[i];
if (!dp[mask].fi) {
cout << -1 << " \n"[i == n - 1];
} else {
cout << dp[mask].se << " \n"[i == n - 1];
}
}
return 0;
} -
对于每个状态,维护一个最大值和一个次大值,再求一个前缀最大值
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38int main() {
std::ios::sync_with_stdio(false);
std::cin.tie(nullptr);
int N;
cin >> N;
int M = 1 << N;
std::vector<int> a(M);
std::vector<std::array<int, 2>> dp(M);
for (int i = 0; i < M; i++) {
cin >> a[i];
dp[i] = {a[i], 0};
}
for (int i = 0; i < N; i++) {
for (int mask = 0; mask < M; mask++) {
if (mask >> i & 1) {
std::array<int, 4> b{dp[mask][0], dp[mask][1], dp[mask ^ (1 << i)][0], dp[mask ^ (1 << i)][1]};
std::sort(b.begin(), b.end(), std::greater());
dp[mask] = {b[0], b[1]};
}
}
}
std::vector<int> F(M);
F[1] = dp[1][0] + dp[1][1];
for (int i = 2; i < M; i++) {
F[i] = std::max(dp[i][0] + dp[i][1], F[i - 1]);
}
for (int i = 1; i < M; i++) {
cout << F[i] << "\n";
}
return 0;
} -
在数组中找 $a_i|(a_j\text{&}a_k),(1\le i<j<k\le n)$ 的最大值
对式子变形,求 $(\overline {a_i}\text{&} (a_j\text{&} a_k))+a_i$ 的最大值,对 $state$ 做超集 SOS-DP,求 $state\subseteq a_j \text{&} a_k$ 的最大下标和次大下标,对 $\overline{a_i}$ 做按位贪心
按位贪心步骤:从高到低枚举每一位 $add$,看 $add$ 当前位是否可以为 1,如果可以 $add\leftarrow add + (1<<j)$,否则跳过枚举下一位,检查即为看 $state=add+(1<<j)$ 中对应的次大下标是否大于 $i$。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48int main() {
std::ios::sync_with_stdio(false);
std::cin.tie(nullptr);
int n;
cin >> n;
std::vector<int> a(n);
int N = 21, M = (1 << N);
std::vector<std::array<int, 2>> dp(M);
auto update = [&](int state, int pos) -> void {
if (pos > dp[state][0]) {
dp[state][1] = dp[state][0];
dp[state][0] = pos;
} else if (pos > dp[state][1]) {
dp[state][1] = pos;
}
};
for (int i = 0; i < n; i++) {
cin >> a[i];
update(a[i], i);
}
for (int i = 0; i < N; i++) {
for (int mask = 0; mask < M; mask++) {
if (!(mask >> i & 1)) {
update(mask, dp[mask ^ (1 << i)][0]);
update(mask, dp[mask ^ (1 << i)][1]);
}
}
}
int ans = 0;
for (int i = 0; i < n - 2; i++) {
int C = (M - 1) ^ a[i], add = 0;
for (int j = N - 1; j >= 0; j--) {
if (!(C >> j & 1) || dp[add ^ (1 << j)][1] <= i) continue;
add += (1 << j);
}
ans = std::max(ans, a[i] + add);
}
cout << ans << "\n";
return 0;
} -
给定 $n$ 个单词,这 $n$ 个单词均为 3 个由
a~x的小写字母组成,规定至少含有一个元音字母的单词是正确的,对于a~x中所有可能的元音字母组合(共有 $2^{24}$ 个),求 $n$ 个单词中正确单词数量的平方和。本质是求 $n$ 个单词中与 $mask$ 字母集合有交集的单词个数,正难则反:把 $mask$ 当作辅音集合,其一定也是
0~(1<<24)-1,对 $dp[mask]$ 做子集和,即为 $n$ 个单词中均为辅音的单词数量 $n-dp[mask]$ 即为正确单词数量。异或和即为答案。1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33int main() {
std::ios::sync_with_stdio(false);
std::cin.tie(nullptr);
int n;
cin >> n;
int N = 24, M = (1 << N);
std::vector<int> dp(M);
for (int i = 0; i < n; i++) {
string s;
cin >> s;
auto [x, y, z] = std::make_tuple(s[0] - 'a', s[1] - 'a', s[2] - 'a');
dp[(1 << x) | (1 << y) | (1 << z)]++;
}
for (int i = 0; i < N; i++) {
for (int mask = 0; mask < M; mask++) {
if (mask >> i & 1) {
dp[mask] += dp[mask ^ (1 << i)];
}
}
}
int ans = 0;
for (int i = 0; i < M; i++) {
ans ^= (n - dp[i]) * (n - dp[i]);
}
cout << ans << "\n";
return 0;
} -
有 $F,G,H$ 三个函数,定义在
0~(1 << N) - 1上,求 $\sum_{x=0}^{2^N-1}R(x),R(x)=\sum_{x\subseteq A\cup B \cup C}F(A)\cdot G(B)\cdot H(C)$考虑求 $T(x)=\sum_{x=A\cup B \cup C}F(A)G(B)H(C)$ 对 $T$ 做子集和DP即可求得 $R$
先对 $F,G,H$ 做子集和DP得 $F’,G’,H’$,表示 $F_i\leftarrow\sum_{j\cup i = i} F_j $ ,设 $a_i = F_i \cdot G_i \cdot H_i$,对 $a$ 做高维前缀和得到 $a_i\leftarrow\sum_{A\cup B\cup C \subseteq i}F(A)G(B)H(C)$
对 $a$ 做高维差分得到 $T_i=\sum_{A\cup B\cup C=i}F(A)G(B)H(C)$,即求得 $T$,再做一遍子集和DP求各项和即可
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57int main() {
std::ios::sync_with_stdio(false);
std::cin.tie(nullptr);
int N, M;
cin >> N;
M = (1 << N);
std::vector<mint> F(M), G(M), H(M);
for (int i = 0; i < M; i++) {
cin >> F[i];
}
for (int i = 0; i < M; i++) {
cin >> G[i];
}
for (int i = 0; i < M; i++) {
cin >> H[i];
}
for (int i = 0; i < N; i++) {
for (int mask = 0; mask < M; mask++) {
if (mask >> i & 1) {
F[mask] += F[mask ^ (1 << i)];
G[mask] += G[mask ^ (1 << i)];
H[mask] += H[mask ^ (1 << i)];
}
}
}
std::vector<mint> a(M);
for (int i = 0; i < M; i++) {
a[i] = F[i] * G[i] * H[i];
}
for (int i = 0; i < N; i++) {
for (int mask = 0; mask < M; mask++) {
if (mask >> i & 1) {
a[mask] -= a[mask ^ (1 << i)];
}
}
}
for (int i = 0; i < N; i++) {
for (int mask = 0; mask < M; mask++) {
if (!(mask >> i & 1)) {
a[mask] += a[mask ^ (1 << i)];
}
}
}
mint ans = std::accumulate(a.begin(), a.end(), mint(0));
cout << ans << "\n";
return 0;
} -
求数组 $a$ 中按位或为 0 的子序列个数
正难则反,考虑有多少子序列按位或为 $mask$
设 $f_i$ 为 $f_i=\sum_{j=1}^n [a_j \text{&} i = i]$
只需要将 $f_i$ 看作桶,记录每个元素出现次数,再做超集和DP即可
设 $g_i=2^{f_i}-1$,即子序列按位或为 $j\land i\subseteq j$ 的个数
考虑将高维前缀和反过来,变成高维差分,对$g$ 即可求得 $g_i$ 为按位或和为 $i$ 的子序列个数
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46int main() {
std::ios::sync_with_stdio(false);
std::cin.tie(nullptr);
int n;
cin >> n;
int N = 20, M = (1 << N);
std::vector<int> f(M);
for (int i = 0; i < n; i++) {
int x;
cin >> x;
f[x]++;
}
for (int i = 0; i < N; i++) {
for (int mask = 0; mask < M; mask++) {
if (!(mask >> i & 1)) {
f[mask] += f[mask ^ (1 << i)];
}
}
}
std::vector<mint> g(M);
for (int mask = 0; mask < M; mask++) {
g[mask] = power(mint(2), f[mask]) - 1;
}
for (int i = 0; i < N; i++) {
for (int mask = 0; mask < M; mask++) {
if (!(mask >> i & 1)) {
g[mask] -= g[mask ^ (1 << i)];
}
}
}
mint ans = power(mint(2), n) - 1;
for (int mask = 1; mask < M; mask++) {
ans -= g[mask];
}
cout << ans << "\n";
return 0;
}

