You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

组合-自然数双向映射规模化优化:非枚举实现方案咨询

解决方案:基于组合数学的无枚举双向映射

核心思路是利用隔板法的组合数性质,直接计算权重元组与ID的对应关系,完全无需预枚举所有组合,内存占用极低,且计算效率不受k大小影响。

关键原理

权重元组本质是满足 $x_1+x_2+...+x_n=k$ 的非负整数解(每个权重为 $x_i/k$),组合总数为 $C(n+k-1, n-1)$。我们通过以下两个核心算法实现双向映射:

  1. Rank算法:给定整数元组,计算其对应的1-based ID
  2. 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.30 19:39:50