简介

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]
$$

解法:

  1. Bruteforce

    1
    2
    3
    4
    5
    6
    7
    for (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)$

  2. 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)$

  3. SoS Dynamic Programming solution

1
2
3
4
5
6
7
8
9
for (int i = 0; i < (1 << N); i++) {
F[i] = A[i];
}
for (int i = 0; i < N; i++) {
for (int mask = 0; mask < (1 << N); mask++) {
if (mask >> i & 1) {
F[mask] += F[mask ^ (1 << 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)];可得高维差分,一般先用高位前缀和,变换后再用高维差分,得到所需的严格等的信息。

例题:

SOS-DP 题单

  1. CF165E

    在数组 $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
    46
    int 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;
    }
  2. arc100_c

    对于每个状态,维护一个最大值和一个次大值,再求一个前缀最大值

    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
    int 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;
    }
  3. CF1208F

    在数组中找 $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
    48
    int 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;
    }
  4. CF383E

    给定 $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
    33
    int 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;
    }
  5. Convering Sets

    有 $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
    57
    int 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;
    }
  6. CF449D

    求数组 $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
    46
    int 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;
    }