数组逐次除以2求和时满足和≥q的排列计数算法问题
问题优化思路与实现方案
暴力枚举全排列的时间复杂度是O(n!),n≥12时计算量就会超过10亿,完全没有实用性,针对这个问题可以根据数据规模选择两类优化方案,核心是避免无意义的重复枚举,同时用动态规划或折半思想合并重复状态。
方案1:状态压缩动态规划(适用n≤20的场景)
核心原理
首先做数值缩放消除浮点误差:题目中的和式为 $S = a_{p1} + \frac{a_{p2}}{2} + \frac{a_{p3}}{4} + ... + \frac{a_{pn}}{2{n-1}}$,要判定$S≥q$,可以把等式两边同时乘以$2{n-1}$,把所有除法转为整数运算,得到等价判定条件:
$$a_{p1}*2^{n-1} + a_{p2}*2^{n-2} + ... + a_{pn}*2^0 ≥ Q$$
其中Q是q缩放后的整数阈值,全程用整数计算可以完全避免浮点数累计精度错误。
我们用二进制数mask表示已经被选中放入排列前k位的元素集合(k为mask中二进制1的个数,即k = __builtin_popcount(mask)),定义dp[mask]为键值对结构,记录选了mask对应元素时,所有可能的累计加权和对应的排列数量。
转移逻辑
- 初始状态:
dp[0] = {0:1},即没有选任何元素时,累计和为0,共1种选法。 - 遍历所有mask,对每个状态,尝试放入所有未被选中的元素i:当前要放的是排列的第k+1位,对应缩放后的权重是$2^{n-1-k}$,新的累计和为原和加上
a[i] * 权重,把计数累加到新maskmask | (1<<i)对应的状态中。 - 最终统计全选状态
dp[(1<<n)-1]中,所有和≥Q的计数之和,就是答案。
优化技巧
- 如果数组元素全为正数,当某个状态的累计和已经≥Q时,后续不管加什么元素(后续元素权重都更小,且为正)最终和一定≥Q,可以直接把这部分计数单独累加,不用再参与后续转移,大幅减少状态量。
- 每个mask对应的和值列表可以提前排序、合并相同和值的计数,减少转移时的遍历次数;如果数组元素范围不大,可以用数组替代哈希表存dp状态,速度会提升数倍。
方案2:折半搜索(Meet-in-the-Middle,适用n≤40的场景)
当n超过20时,状压DP需要的2^20以上的内存和计算量会开始吃紧,这时可以用折半思想把问题规模拆半:
- 把n个元素平均分成两组,每组最多20个元素,分别用上述状压方法预处理出每组的状态:即选k个元素时,所有可能的排列和值对应的计数。
- 枚举前半部分排列的长度k,以及k个元素中来自第一组、第二组的元素个数,把前半部分所有可能的和值存下来,后半部分所有可能的和值排序后预处理前缀和。
- 对每个前半部分的和s,通过二分查找快速统计有多少个后半部分的和t满足
s + t/(2^k) ≥ q,把对应的排列数相乘累加即可得到总结果。
折半的时间复杂度是O(n*2{n/2}),n=40时220约为100万,普通家用电脑就可以快速算出结果。
参考实现(状压DP版本)
#include <bits/stdc++.h> using namespace std; typedef long long ll; ll calc_valid(vector<int>& arr, ll threshold, int n) { vector<unordered_map<ll, ll>> dp(1 << n); dp[0][0] = 1; for (int mask = 0; mask < (1 << n); ++mask) { int selected = __builtin_popcount(mask); ll cur_weight = 1LL << (n - 1 - selected); for (auto& [sum, cnt] : dp[mask]) { // 正数剪枝:当前和已经超过阈值,剩下的位置不管填什么正数都满足条件 if (sum >= threshold) { // 剩下的n-selected个元素全排列,直接加计数,不用继续转移 ll rest = 1; for (int i = 1; i <= n - selected; ++i) rest *= i; dp[(1<<n)-1][threshold] += cnt * rest; continue; } for (int i = 0; i < n; ++i) { if (!(mask & (1 << i))) { ll new_sum = sum + 1LL * arr[i] * cur_weight; dp[mask | (1 << i)][new_sum] += cnt; } } } } ll ans = 0; for (auto& [sum, cnt] : dp[(1 << n) - 1]) { if (sum >= threshold) ans += cnt; } return ans; } int main() { int n; double q; scanf("%d %lf", &n, &q); vector<int> arr(n); for (int i = 0; i < n; ++i) scanf("%d", &arr[i]); // 缩放阈值为整数,注意如果q有小数要做 rounding 避免精度损失 ll q_scaled = llround(q * (1LL << (n-1))); printf("%lld\n", calc_valid(arr, q_scaled, n)); return 0; }
内容的提问来源于stack exchange,提问作者dnaiel
相关产品推荐
相关产品推荐

