组合-自然数双向映射规模化优化:非枚举实现方案咨询
解决方案:基于组合数学的无枚举双向映射
核心思路是利用隔板法的组合数性质,直接计算权重元组与ID的对应关系,完全无需预枚举所有组合,内存占用极低,且计算效率不受k大小影响。
关键原理
权重元组本质是满足 $x_1+x_2+...+x_n=k$ 的非负整数解(每个权重为 $x_i/k$),组合总数为 $C(n+k-1, n-1)$。我们通过以下两个核心算法实现双向映射:
- Rank算法:给定整数元组,计算其对应的1-based ID
- Unrank算法:给定ID,反推对应的整数元组
组合数计算工具
首先实现高效的组合数计算,并用缓存避免重复计算:
from functools import lru_cache @lru_cache(maxsize=None) def comb(a, b): if b < 0 or b > a: return 0 if b == 0 or b == a: return 1 b = min(b, a - b) # 优化计算量 result = 1 for i in range(1, b + 1): result = result * (a - b + i) // i return result
Rank算法(元组转ID)
通过组合数累加公式,直接计算当前元组在所有组合中的排名,无需遍历:
def rank_tuple(int_tuple, k, n): """将整数元组转换为1-based ID""" if n == 1: return 1 x_n = int_tuple[-1] # 计算所有x_n' < 当前x_n的组合总数 total_prev = comb(k + n - 1, n - 1) - comb((k - x_n) + n - 1, n - 1) # 递归计算前n-1个元素的子排名 sub_tuple = int_tuple[:-1] sub_rank = _rank_sub_tuple(sub_tuple, k - x_n, n - 1) return total_prev + sub_rank + 1 def _rank_sub_tuple(sub_tuple, s, n): """计算子元组在当前约束下的0-based排名""" if n == 2: x1, _ = sub_tuple return s - x1 x1 = sub_tuple[0] # 所有x1' > 当前x1的组合总数 cnt_prev = comb((s - x1) + n - 2, n - 2) # 递归计算剩余元素的子排名 sub_sub_tuple = sub_tuple[1:] sub_sub_rank = _rank_sub_tuple(sub_sub_tuple, s - x1, n - 1) return cnt_prev + sub_sub_rank
Unrank算法(ID转元组)
通过二分查找快速定位元组各维度的值,避免线性遍历:
def unrank_id(idx, k, n): """将1-based ID转换为整数元组""" idx_0 = idx - 1 # 转为0-based索引 if n == 1: return (k,) total_comb = comb(k + n - 1, n - 1) # 二分查找确定最后一个元素的值 left, right = 0, k best_m = 0 while left <= right: mid = (left + right) // 2 current_sum = total_comb - comb((k - mid) + n - 1, n - 1) if current_sum <= idx_0: best_m = mid left = mid + 1 else: right = mid - 1 m = best_m current_sum = total_comb - comb((k - m) + n - 1, n - 1) sub_idx = idx_0 - current_sum # 递归计算前n-1个元素 sub_tuple = _unrank_sub_idx(sub_idx, k - m, n - 1) return sub_tuple + (m,) def _unrank_sub_idx(sub_idx, s, n): """根据子索引反推子元组""" if n == 2: x1 = s - sub_idx return (x1, sub_idx) # 二分查找确定第一个元素的值 left, right = 0, s best_x1_rev = 0 total = 0 while left <= right: mid = (left + right) // 2 c = comb(mid + n - 3, n - 3) if total + c <= sub_idx: total += c best_x1_rev = mid left = mid + 1 else: right = mid - 1 x1 = s - best_x1_rev new_sub_idx = sub_idx - total sub_sub_tuple = _unrank_sub_idx(new_sub_idx, s - x1, n - 1) return (x1,) + sub_sub_tuple
封装为WeightSet类
保持与原代码一致的接口,用自定义映射类实现动态计算:
class WeightTupleFromID: """ID到权重元组的映射类""" def __init__(self, k, n): self.k = k self.n = n def __getitem__(self, idx): int_tuple = unrank_id(idx, self.k, self.n) return tuple(x / self.k for x in int_tuple) def __contains__(self, idx): return 1 <= idx <= comb(self.k + self.n - 1, self.n - 1) class WeightIDFromTuple: """权重元组到ID的映射类""" def __init__(self, k, n): self.k = k self.n = n def __getitem__(self, weight_tuple): # 处理浮点数精度问题,转为整数元组 int_tuple = tuple(round(x * self.k) for x in weight_tuple) assert sum(int_tuple) == self.k, "无效的权重元组" return rank_tuple(int_tuple, self.k, self.n) def __contains__(self, weight_tuple): try: int_tuple = tuple(round(x * self.k) for x in weight_tuple) return sum(int_tuple) == self.k and all(x >= 0 for x in int_tuple) except: return False class WeightSet: def __init__(self, k, n): self._k = k self._n = n self._weight_tuple_cnt = comb(k + n - 1, n - 1) @property def n(self): return self._n @property def k(self): return self._k @property def weight_tuple_cnt(self): return self._weight_tuple_cnt @property def weight_tuple_from_id(self): return WeightTupleFromID(self._k, self._n) @property def weight_id_from_tuple(self): return WeightIDFromTuple(self._k, self._n)
测试验证
与原代码示例结果完全一致:
weight_set1 = WeightSet(10,3) print(weight_set1.weight_tuple_cnt) # 输出:66 print(weight_set1.weight_tuple_from_id[20]) # 输出:(0.1, 0.8, 0.1) print(weight_set1.weight_id_from_tuple[(0.1,0.8,0.1)]) # 输出:20
优势说明
- 内存高效:无需存储所有组合,内存占用仅为组合数缓存,可轻松处理k≥100的场景
- 计算快速:每次查询时间复杂度为O(n log k),远快于枚举法
- 接口兼容:完全保留原代码的调用方式,无需修改业务逻辑
内容的提问来源于stack exchange,提问作者Dino Hsu
相关产品推荐
相关产品推荐

