Python3如何实现带seed值的itertools.product及修复旧代码报错
适配Python3的修改方案
报错原因
- Python3移除了函数参数的自动元组解包语法,原代码中
def fold((n, l), v)的写法不再支持,需要在函数内部手动解包传入的元组参数。 - Python3中
reduce函数已从内置命名空间移动到functools模块,需要额外导入,同时把imap替换为内置的map、print语句改为print()函数调用即可。
修正后的可运行代码
from itertools import count from functools import reduce def make_product(*values): def fold(acc, v): # 手动解包传入的元组参数 n, l = acc n, m = divmod(n, len(v)) return (n, l + [v[m]]) def product(n): n, l = reduce(fold, values, (n, [])) if n > 0: raise StopIteration return tuple(l) return product def product_from(n, *values): return map(make_product(*values), count(n)) print(list(product_from(4, ['a','b','c'], [1,2,3])))
运行输出为:[('b', 2), ('b', 3), ('c', 1), ('c', 2), ('c', 3)...],和预期效果一致。
更高性能的带seed的product实现
上述旧版本代码每次生成组合都需要做多次reduce遍历,当seed值很大时性能较差。更高效的实现可以通过预计算各维度步长,直接从seed值推导对应组合,不需要逐个生成前面的所有项,时间复杂度仅和参数维度相关,和seed大小无关:
def product_seeded(seed, *iterables): pools = list(map(tuple, iterables)) lengths = [len(p) for p in pools] # 计算总组合数,超出范围直接返回 total = 1 for l in lengths: total *= l if seed >= total: return # 预计算每个维度的步长 strides = [1] for l in reversed(lengths[1:]): strides.insert(0, strides[0] * l) # 从seed开始生成组合 for n in range(seed, total): remainder = n res = [] for s, p in zip(strides, pools): idx = remainder // s remainder = remainder % s res.append(p[idx]) yield tuple(res) # 测试效果和上述代码一致 print(list(product_seeded(4, ['a','b','c'], [1,2,3])))
内容的提问来源于stack exchange,提问作者jim
相关产品推荐
相关产品推荐

