容斥原理不是一套需要背的数学公式。把它当成一个工程问题就够了:多个规则都会产出同一批对象时,如何只统计一次。

例如,三个上游服务都上报了用户 ID。把三个服务的数量直接相加,会把同时出现于多个服务的用户重复计算。容斥做的事就是:

  1. 先把每个来源的数量相加;
  2. 把两个来源共同产出的对象减掉;
  3. 如果一个对象同时属于三个来源,它被减多了,再加回来。

本文用 LeetCode 3116 说明这个过程。重点不是记住公式,而是识别这种数据形状:规则少、候选范围大、重叠很难显式去重,但任意一组规则的重叠数量很好算。

先看一个不抽象的例子

要统计 111212 中,能被 23 整除的整数数量。

可以把两个规则看作两个“产出流”:

  • 2 的流:2, 4, 6, 8, 10, 12,共 66 个;
  • 3 的流:3, 6, 9, 12,共 44 个。

直接相加得到 1010,但这是错的:612 都来自两个流,各数了两次。

两个流共同产出的数,必须同时是 23 的倍数,也就是 6 的倍数:6, 12,共 22 个。

所以唯一对象数是:

6+42=86 + 4 - 2 = 8

实际集合为 {2, 3, 4, 6, 8, 9, 10, 12},正好 88 个。

容斥 = 合并多个有重复项的逻辑数据源时,按交集回冲重复计数。

三个规则时为什么又要加回来

假设某个对象同时命中规则 ABC

步骤这个对象被计数几次
加三个单独规则+3+3
减三个两两重叠3-3
加三个规则的共同重叠+1+1
最终11

这就是“加、减、加”的原因。一个对象若只命中两个规则,则只经历 +21=1+2 - 1 = 1。无论它命中几条规则,最后都会保留一次。

因此固定规则只有一句:

  • 选中奇数个规则:加上它们共同产出的数量;
  • 选中偶数个规则:减去它们共同产出的数量。

不用先理解二项式定理;先把它理解成一个可靠的去重计数协议

工程上什么时候该想到它

容斥不是“集合题通用解”。它需要同时满足以下条件:

条件工程含义
规则数量小要枚举所有规则组合,成本是 2n2^n;通常 n20n \le 20 才安全。
候选范围大不能把所有对象逐个生成后放进 set 去重。
交集可快速计数给定一组规则,必须能直接算出“同时命中它们”的数量。
只需要数量或单调判定它擅长 count(x),不擅长列出所有对象。

常见信号:n <= 15,但值域、答案或第 kk 项达到 10910^9 甚至更大。这时“枚举对象”通常不可能,而“枚举规则子集”反而很便宜。

在 3116 里,集合到底是什么

题目给若干面额 coins。每次只能选一种面额、可使用任意张,金额 c 能产生:

1
c, 2c, 3c, 4c, ...

例如 coins = {2, 3}

1
2
面额 2 产生:2, 4, 6, 8, 10, 12, ...
面额 3 产生:3, 6, 9, 12, ...

612 同时存在于两个流。题目要的是不同金额的第 kk 小,而不是“带来源标签的金额”第 kk 小,所以不能重复计数。

定义:

count(x)=不超过 x 的不同可组成金额数count(x) = \text{不超过 }x\text{ 的不同可组成金额数}

对单个面额 c,不超过 x 的产出数是:

xc\left\lfloor \frac{x}{c} \right\rfloor

对多个面额,交集也能直接算。一个金额同时在面额 ab 的流里,等价于它同时是 ab 的倍数,等价于它是 lcm(a, b) 的倍数:

AaAb=xlcm(a,b)\left|A_a \cap A_b\right| = \left\lfloor \frac{x}{\operatorname{lcm}(a, b)} \right\rfloor

例如 a = 2b = 3,最小公倍数是 6。所以:

count(12)=122+123126=6+42=8count(12) = \left\lfloor\frac{12}{2}\right\rfloor + \left\lfloor\frac{12}{3}\right\rfloor - \left\lfloor\frac{12}{6}\right\rfloor = 6 + 4 - 2 = 8

这就是本题的完整建模:每种面额是一条无限金额流;lcm 告诉我们多条流的重叠频率;容斥给出合并后的去重数量。

为什么要二分,而不是生成前 k 个金额

count(x) 有一个关键性质:x 变大时它不会变小。

1
2
x:          1  2  3  4  5  6  7  8  9 10 11 12
count(x): 0 1 2 3 3 4 4 5 6 6 6 7

因此第 kk 小金额,就是第一个满足 count(x) >= kx。这正是二分答案的标准模型:

1
2
3
猜 x
├─ count(x) >= k:答案在左边,保留 x
└─ count(x) < k:x 太小,去右边

上界可取 min(coins) * k。只使用最小面额,也能得到 minCoin, 2*minCoin, ..., k*minCoinkk 个不同金额,所以第 kk 小一定不会更大。

可直接提交的实现

代码分两段:

  1. 预处理每个面额组合的 lcm 和容斥符号;
  2. 二分时扫描这些组合,得到 count(x)
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
58
59
60
61
#include <algorithm>
#include <numeric>
#include <vector>

class Solution {
public:
using int64 = long long;

long long findKthSmallest(std::vector<int>& coins, int k) {
const int n = static_cast<int>(coins.size());
const int state_count = 1 << n;
const int64 min_coin = *std::min_element(coins.begin(), coins.end());
const int64 hi = min_coin * static_cast<int64>(k);

// lcm[mask]:mask 选中的面额共同出现时的周期。
// odd[mask]:选中面额数是否为奇数;奇数加,偶数减。
std::vector<int64> lcm(state_count, 1);
std::vector<char> odd(state_count, false);

for (int mask = 1; mask < state_count; ++mask) {
const int bit = __builtin_ctz(static_cast<unsigned>(mask));
const int previous = mask & (mask - 1); // 去掉最低位的已选面额
odd[mask] = !odd[previous];

// 已经超过二分上界的 lcm 永远没有贡献,不再做乘法。
if (lcm[previous] > hi) {
lcm[mask] = hi + 1;
continue;
}

const int64 coin = coins[bit];
const int64 base = lcm[previous] / std::gcd(lcm[previous], coin);
// 先检查再乘;lcm 大于 hi 时只保留“无贡献”标记,避免溢出。
lcm[mask] = base > hi / coin ? hi + 1 : base * coin;
}

const auto count_at_most = [&](int64 x) {
int64 result = 0;
for (int mask = 1; mask < state_count; ++mask) {
if (lcm[mask] > x) {
continue;
}
const int64 overlap_count = x / lcm[mask];
result += odd[mask] ? overlap_count : -overlap_count;
}
return result;
};

int64 lo = 1;
int64 right = hi;
while (lo < right) {
const int64 mid = lo + (right - lo) / 2;
if (count_at_most(mid) >= k) {
right = mid; // mid 已经包含至少 k 个金额,继续找更小的。
} else {
lo = mid + 1;
}
}
return lo;
}
};

代码中的三个关键防线

  1. 不生成金额k 最大可达 2×1092 \times 10^9。小根堆或 set 需要按金额个数工作,规模不成立;这里每次只算数量。
  2. long long 和先除后乘lcm 可能超过 int。计算 lcm(a, b) 时使用 a / gcd(a, b) * b,并在乘法前确认不会超过二分上界。
  3. lcm 超过 hi 立即截断:对任何二分中的 xhix \le hix/lcm=0\lfloor x / lcm \rfloor = 0。把它记成 hi + 1 就足够了,也避免无意义计算。

复杂度:为什么它能过

设面额数为 nn

  • 子集数:2n12^n - 1
  • 每次 count_at_most(x)O(2n)O(2^n)
  • 二分次数:O(log(minCoink))O(\log(minCoin \cdot k))
  • 总复杂度:O(2nlog(minCoink))O(2^n \log(minCoin \cdot k))

本题 n15n \le 15,最多只有 3276732767 个非空子集;二分约 3636 次。核心工作量约一百万次整数运算,远小于按 kk 生成金额。

记忆方式

遇到题目时,不要先问“能不能套容斥公式”,而是按这个顺序检查:

  1. 不同规则是否会产出同一个对象? 会,存在去重问题。
  2. 对象范围是否大到不能逐个生成? 会,不能用枚举加 set
  3. 任意几条规则同时命中的数量能否直接算? 能,容斥可用。
  4. 规则数量是否小? 是,能枚举所有规则组合。
  5. 目标是否是第 kk 小或阈值? 是,把容斥得到的 count(x) 放进二分。

把容斥看成“不落地生成数据的去重计数器”,就不容易混淆:它不负责找元素,也不负责去重存储;它只用交集数量,算出合并结果中的唯一对象数。