【问题标题】:How can implement efficiently this subset enumeration problem?如何有效地实现这个子集枚举问题?
【发布时间】:2020-09-08 21:39:20
【问题描述】:

我有一个数字列表Sn = [a, b, c, d, ...] 和一组不重叠的间隔Si = {I1, I2, I3, ...}。鉴于此,我的问题是找到Sn 子集的列表L,使得每个子集中的元素总和至少绑定到Si 内的一个区间内。

我现在的方法是枚举Sn 的所有子集,并根据它们是否适合一个区间来过滤它们。这是正确的,但效率低。

import itertools

Sn = [12, 30, 60, 6, 6]
Si = {(12, 12), (18, 24), (30, 48)}

def enumerate_sets():
    sets = []
    for i in range(len(Sn)):
        for comb in itertools.combinations(Sn, i + 1):
            for interval in Si:
                if interval[0] <= sum(comb) <= interval[1]:
                    sets.append(comb)
                    break
    return sets

print(enumerate_sets())
# [(12,), (30,), (12, 30), (12, 6), (12, 6), (30, 6), (30, 6), (6, 6), (12, 30, 6), (12, 30, 6), (12, 6, 6), (30, 6, 6)]

如何有效地实现这个子集问题?首选 python 中的答案,但任何(伪)语言都可以。

【问题讨论】:

  • 您使用什么指标来衡量程序的“效率”?算法复杂度?内存消耗?实际运行时间?
  • @ethane 算法复杂性,虽然最终它的运行时间很重要,但我没有足够详尽的样本上关于运行时间的数据
  • 请添加您的解决方案,以便我们了解您的解决方案的复杂性并更好地回答问题。
  • 只是为了澄清-问题是找到Sn的所有子集的列表L?或者你想要一些特定的数量?
  • 一组数字对你来说意味着一组唯一的数字,对吧?

标签: python algorithm


【解决方案1】:

由于组合的总和足以知道它是否是一个有效的组合,你可以将它作为一个设置键,避免做额外的工作。

enumerate_sets_2 基准比这里的原始基准快(超过)2 倍:

import timeit
import itertools
import collections

Sn = [12, 30, 60, 6, 6, 92, 443, -8, 112, 96]
Si = {(12, 12), (18, 24), (30, 48)}


def enumerate_sets_orig(Sn, Si):
    sets = []
    for i in range(len(Sn)):
        for comb in itertools.combinations(Sn, i + 1):
            for interval in Si:
                if interval[0] <= sum(comb) <= interval[1]:
                    sets.append(comb)
                    break
    return sets


def enumerate_sets_2_iter(Sn, Si):
    valid_sums = set()
    for i in range(len(Sn)):
        for comb in itertools.combinations(Sn, i + 1):
            comb_sum = sum(comb)
            if comb_sum in valid_sums:
                yield comb
                continue
            for a, b in Si:
                if a <= comb_sum <= b:
                    valid_sums.add(comb_sum)
                    yield comb
                    break



def make_comparable(result):
    # Make an enumerate_sets_* result comparable by sorting and removing duplicates
    return set(tuple(sorted(int(v) for v in i)) for i in result)


def t(f):
    # Wrap the function to evaluate generators and to pass in the args
    fw = lambda: list(f(Sn, Si))
    # Validate this solution before benchmarking
    assert expected == make_comparable(fw())
    # Benchmark
    count, time_taken = timeit.Timer(fw).autorange()
    print(f"{f.__name__:<25} {count / time_taken:>10.2f} iter/s")

expected = make_comparable(enumerate_sets_orig(Sn, Si))

t(enumerate_sets_orig)
t(enumerate_sets_2_iter)

【讨论】:

    【解决方案2】:

    好的,至于编程,python 解决方案 - 这可以解决问题:

    from itertools import combinations
    import numpy as np
    
    def getL(Sn:np.ndarray, Si:np.ndarray):
        Si.sort(axis=0)
        Sn.sort()
        for i in range(1, len(Sn)+1):
            for s in combinations(Sn, i):
                su=sum(s)
                if((np.logical_and(Si[:,0]<=su, Si[:,1]>=su)).any()):
                    yield s
    

    样本输出:

    >>> # to convert your input into numpy:
    >>> Sn = [12, 30, 60, 6, 6]
    >>> Si = {(12, 12), (18, 24), (30, 48)}
    >>> Sn = np.array(Sn)
    >>> Si = np.array(list(Si))
    >>> print(list(getL(Sn, Si)))
    
    [(12,), (30,), (6, 6), (6, 12), (6, 30), (6, 12), (6, 30), (12, 30), (6, 6, 12), (6, 6, 30), (6, 12, 30), (6, 12, 30)]
    

    很少注意-为了节省内存-通常使用generator应该更快,因此您不会只是累积到list,但是一旦您得到一些东西-您会返回,然后忘记它。 使用numpy - 这将显着加快浏览间隔。

    【讨论】:

    • OP 的 SnSi 不是 np.ndarrays... 你能用 OP 的数据显示这个用法吗?
    • 由于堆栈开销,生成器实际上通常比函数慢。在OP的情况下是否重要是另一回事:)
    • 你说得对,如果你有足够的内存,并且需要一次存储整个东西,这种东西很容易增长——这就是我建议生成器的方式。此外,如果您正在尝试找到最合适的 - 您只需将最好的保留到现在,将其与下一个进行比较,并保持更合适的。你不必存储整个东西;)
    • 奇怪的是,尽管 Numpy 通常很快,但这个基准测试比我下面的解决方案慢得多(当然,当你实际迭代它时)。
    • 不,我的意思是假设N个元素不匹配意味着N+1个元素不匹配;如果第 N+1 个元素减少总和怎么办?
    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 2013-01-10
    • 2011-12-11
    • 1970-01-01
    • 2023-03-13
    • 1970-01-01
    • 2023-04-03
    • 1970-01-01
    相关资源
    最近更新 更多