问题引入
想象这样一个场景:给定一组硬币面值 coins = [2, 3, 5],我们需要找到第 k 小的能被至少一个硬币面值整除的正整数。
例如,前几个这样的数是:2, 3, 4, 5, 6, 8, 9, 10…
第5个是6,第8个是10。
这个问题看似简单,但当 k 很大(如10^9)或硬币种类很多时,直接枚举就不可行了。今天我们来探讨一个巧妙的解决方案。
核心思路
这个问题的解决方案由两个核心思想组成:
1. 二分查找(Binary Search)
如果我们可以快速计算 "1到x之间有多少个合法数",那么就可以用二分查找找到第k个数。
2. 容斥原理(Inclusion-Exclusion Principle)
用于精确计算1到x之间合法数的个数,避免重复计数。
容斥原理解析
什么是容斥原理?//力扣2652练习
容斥原理是组合数学中计算多个集合并集大小的公式:
text
|A ∪ B ∪ C| = |A| + |B| + |C| – |A∩B| – |A∩C| – |B∩C| + |A∩B∩C|
应用于我们的问题
设 Aᵢ = {能被coins[i]整除的数},我们需要计算:
text
|A₁ ∪ A₂ ∪ … ∪ Aₙ|
关键观察:
-
一个数能被多个硬币同时整除 ⇔ 能被它们的最小公倍数(LCM)整除
-
例如:能被2和3同时整除 ⇔ 能被6整除
计算1到x之间的合法数
对于每个非空子集,计算 ⌊x / lcm(子集)⌋,这是1到x之间能被该子集所有硬币整除的数的个数。
然后根据子集大小的奇偶性,决定加还是减:
-
奇数个元素:加
-
偶数个元素:减
算法设计
数据结构准备
cpp
int n = coins.size();
int m = 1 << n; // 2^n个子集
vector<int> bit_count(m); // 每个子集中硬币数量
vector<ll> lcm(m); // 每个子集的LCM(存储所有子集最小公倍数)
1. 预处理所有子集的LCM
cpp
for (int mask = 1; mask < m; mask++) {
ll cur_lcm = 1;
for (int i = 0; i < n; i++) {
if (mask >> i & 1) {//通过左移判断该位是否为1,比如mask=2=0010,i=1,(mask>>i=0001)&1=1,存在
// 计算 LCM,防止溢出–>一般通过除法防止溢出
ll tmp = cur_lcm / gcd(cur_lcm, coins[i]);//gcd获得最大公约数
if (tmp <= r / coins[i]) {
cur_lcm = tmp * (ll)coins[i];
} else {
cur_lcm = r + 1; // 标记为超出范围
break;
}
bit_count[mask]++;//选中的数—+1
}
}
lcm[mask] = cur_lcm;//最小公约数
}
为什么需要防溢出?
LCM可能非常大,超出long long范围,所以我们需要检查乘法是否会溢出。
2. 计算合法数个数
cpp
auto count_valid = [&](ll x) -> ll {//Lambda表达式(匿名函数)语法
| auto get | 定义变量 | get是一个变量,类型由编译器自动推断 |
| = | 赋值 | 将Lambda表达式赋值给get |
| [&] | 捕获列表 | 表示按引用捕获外部所有变量 |
| (ll x) | 参数列表 | 这个Lambda接受一个ll类型的参数x |
| -> ll | 返回类型 | 指定返回类型为ll(long long) |
| { … } | 函数体 | Lambda的执行代码 |
ll count = 0;
for (int mask = 1; mask < m; mask++) {//每次都遍历所有子集
if (lcm[mask] > x) continue;//超过最大,跳过
if (bit_count[mask] & 1) {
count += x / lcm[mask]; // 奇数个元素:加
} else {
count -= x / lcm[mask]; // 偶数个元素:减
}
}
return count;
};
3. 二分查找
cpp
ll left = k, right = 1ll * coins[0] * k + 1;
while (left < right) {
ll mid = (left + right) >> 1;//左移也是/2
if (count_valid(mid) >= k) {
right = mid; // 答案在左半部分
} else {
left = mid + 1; // 答案在右半部分
}
}
return left;
完整代码//出自力扣解析
class Solution {
public:
using ll = long long;//别名 ll = long long
long long findKthSmallest(vector<int>& coins, int k) {
int n = coins.size();
int m = (1 << n);//看有多少个子集组合
sort(coins.begin(), coins.end());
vector<int> bit_count(m);//每个子集(mask)中包含的硬币数量
vector<ll> lcm(m);//存储所有子集最小公倍数(LCM)的数组
ll l = k, r = ll(coins[0] ) * k + 1;//通过题意分析最大与最小值
for (int mask = 1; mask < m; mask++) {
ll cur_lcm = 1;
for (int i = 0; i < n; i++) {
if (mask >> i & 1) {
ll tmp = cur_lcm / gcd(cur_lcm, coins[i]);//GCD 是 Greatest Common Divisor 的缩写,中文意思是最大公约数(也叫最大公因数)
if (tmp <= r / coins[i]) {
cur_lcm = tmp * coins[i];//最小公倍数
} else {
cur_lcm = r + 1;
break;
}
bit_count[mask]++;
}
}
lcm[mask] = cur_lcm;
}
auto get = [&](ll x) -> ll {
ll count = 0;
for (int mask = 1; mask < m; mask++) {
if (lcm[mask] > x) {
continue;
}
if (bit_count[mask] & 1) {
count += x / lcm[mask];
} else {
count -= x / lcm[mask];
}
}
return count;
};
while (l < r) {//二分查找更快
ll x = (l + r) >> 1;
if (get(x) >= k) {
r = x;
} else {
l = x + 1;
}
}
return l;
}
};
优化建议:对于 n 较大的情况,可以考虑:
去重:去除能被其他硬币整除的硬币(如6可以被2和3替代)
最小堆法:使用优先队列生成前k个数
莫比乌斯反演:更数学化的处理方式
例如在排序后去重:
vector<int> remove_redundant_gcd(vector<int>& coins) {
sort(coins.begin(), coins.end()); // ① 排序
vector<int> filtered; // ② 存储保留的硬币
for (int i = 0; i < coins.size(); i++) { // ③ 遍历每个硬币
bool is_redundant = false; // ④ 标记是否冗余
for (int j = 0; j < filtered.size(); j++) { // ⑤ 与已保留的硬币比较
if (coins[i] % filtered[j] == 0) { // ⑥ 判断是否能被整除
is_redundant = true;
break; // ⑦ 发现冗余,立即跳出
}
}
if (!is_redundant) { // ⑧ 如果不是冗余的
filtered.push_back(coins[i]); // ⑨ 保留该硬币
}
}
return filtered;
}
如有错误,欢迎在评论区交流讨论指正。
