【问题标题】:Python Alternative to itertools product with numpyPython 用 numpy 替代 itertools 产品
【发布时间】:2018-03-25 12:00:36
【问题描述】:

我正在使用大小不一的列表列表。例如,alternativesList 可以在一次迭代中包含 4 个列表,在另一次迭代中包含 7 个列表。

我想要做的是捕获不同列表中的每个单词组合。

让我们这么说

a= [1,2,3]
alternativesList.append(a)
b = ["a","b","c"]
alternativesList.append(b)

productList = itertools.product(*alternativesList)

将创造

[(1, 'a'), (1, 'b'), (1, 'c'), (2, 'a'), (2, 'b'), (2, 'c' ), (3, 'a'), (3, 'b'), (3, 'c')]

这里的一个问题是我的 productList 可能太大,可能会导致内存问题。所以我使用 productList 作为对象并稍后对其进行迭代。

我想知道的是,有没有办法用 numpy 创建相同的对象,它比 itertools 工作得更快?

【问题讨论】:

标签: python python-3.x numpy itertools


【解决方案1】:

一般来说,如果我们将优化视为一个天平,内存和运行时间将是它的两个秤盘。这就是说内存优化和运行时优化之间存在间接关系(并非总是如此,但在大多数情况下)。现在,关于您的问题:

有没有办法用 numpy 创建比 itertools 更快的相同对象?

肯定有,但您需要注意的另一点是抽象会给您更大的灵活性,而这正是 itertools.product 给您的,而 Numpy 没有。如果在这种情况下可伸缩性不是一个重要的因素,您可以使用 Numpy 做到这一点,并且不会放弃任何好处。这是使用column_stackrepeattile 函数的一种方法:

In [5]: np.column_stack((np.repeat(a, b.size),np.tile(b, a.size)))
Out[5]: 
array([['1', 'a'],
       ['1', 'b'],
       ['1', 'c'],
       ['2', 'a'],
       ['2', 'b'],
       ['2', 'c'],
       ['3', 'a'],
       ['3', 'b'],
       ['3', 'c']], dtype='<U21')

现在,仍然有一些方法可以通过使用 U2U1 等更轻的类型来使该数组占用更少的内存。

In [10]: np.column_stack((np.repeat(a, b.size),np.tile(b, a.size))).astype('U1')
Out[10]: 
array([['1', 'a'],
       ['1', 'b'],
       ['1', 'c'],
       ['2', 'a'],
       ['2', 'b'],
       ['2', 'c'],
       ['3', 'a'],
       ['3', 'b'],
       ['3', 'c']], dtype='<U1') 

【讨论】:

  • 不幸的是,可扩展性对我来说是一个重要因素。在你回答之后,我做了一些研究,发现这对我有用:numpy.array(numpy.meshgrid(*alternativesList)).T.reshape(-1,len(alternativesList))。但由于某种原因,它的工作速度较慢。而且看起来它会造成内存问题
【解决方案2】:

您可以通过显式指定复合 dtype 来避免 numpy 尝试查找包罗万象的 dtype 所引起的一些问题:

代码+一些时间安排:

import numpy as np
import itertools

def cartesian_product_mixed_type(*arrays):
    arrays = *map(np.asanyarray, arrays),
    dtype = np.dtype([(f'f{i}', a.dtype) for i, a in enumerate(arrays)])
    out = np.empty((*map(len, arrays),), dtype)
    idx = slice(None), *itertools.repeat(None, len(arrays) - 1)
    for i, a in enumerate(arrays):
        out[f'f{i}'] = a[idx[:len(arrays) - i]]
    return out.ravel()

a = np.arange(4)
b = np.arange(*map(ord, ('A', 'D')), dtype=np.int32).view('U1')
c = np.arange(2.)

np.set_printoptions(threshold=10)

print(f'a={a}')
print(f'b={b}')
print(f'c={c}')

print('itertools')
print(list(itertools.product(a,b,c)))
print('numpy')
print(cartesian_product_mixed_type(a,b,c))

a = np.arange(100)
b = np.arange(*map(ord, ('A', 'z')), dtype=np.int32).view('U1')
c = np.arange(20.)

import timeit
kwds = dict(globals=globals(), number=1000)

print()
print(f'a={a}')
print(f'b={b}')
print(f'c={c}')

print(f"itertools: {timeit.timeit('list(itertools.product(a,b,c))', **kwds):7.4f} ms")
print(f"numpy:     {timeit.timeit('cartesian_product_mixed_type(a,b,c)', **kwds):7.4f} ms")

a = np.arange(1000)
b = np.arange(1000, dtype=np.int32).view('U1')

print()
print(f'a={a}')
print(f'b={b}')

print(f"itertools: {timeit.timeit('list(itertools.product(a,b))', **kwds):7.4f} ms")
print(f"numpy:     {timeit.timeit('cartesian_product_mixed_type(a,b)', **kwds):7.4f} ms")

样本输出:

a=[0 1 2 3]
b=['A' 'B' 'C']
c=[0. 1.]
itertools
[(0, 'A', 0.0), (0, 'A', 1.0), (0, 'B', 0.0), (0, 'B', 1.0), (0, 'C', 0.0), (0, 'C', 1.0), (1, 'A', 0.0), (1, 'A', 1.0), (1, 'B', 0.0), (1, 'B', 1.0), (1, 'C', 0.0), (1, 'C', 1.0), (2, 'A', 0.0), (2, 'A', 1.0), (2, 'B', 0.0), (2, 'B', 1.0), (2, 'C', 0.0), (2, 'C', 1.0), (3, 'A', 0.0), (3, 'A', 1.0), (3, 'B', 0.0), (3, 'B', 1.0), (3, 'C', 0.0), (3, 'C', 1.0)]
numpy
[(0, 'A', 0.) (0, 'A', 1.) (0, 'B', 0.) ... (3, 'B', 1.) (3, 'C', 0.)
 (3, 'C', 1.)]

a=[ 0  1  2 ... 97 98 99]
b=['A' 'B' 'C' ... 'w' 'x' 'y']
c=[ 0.  1.  2. ... 17. 18. 19.]
itertools:  7.4339 ms
numpy:      1.5701 ms

a=[  0   1   2 ... 997 998 999]
b=['' '\x01' '\x02' ... 'ϥ' 'Ϧ' 'ϧ']
itertools: 62.6357 ms
numpy:      8.0249 ms

【讨论】:

  • 太棒了。它做我想做的事情。但是,我相信它首先将所有可能的组合放入列表而不是返回它。这可能会给我带来记忆问题。
  • @rkd 它将所有内容存储在一个大数组中是的。在最后一个示例中,您从中获得的是 ~6 的因子:一个对列表每对占用 72 个字节 - 列表中指向元组的指针为 8 个,元组为 48 个,指向元组的两个指针为 16 个元素。该数组直接存储所有内容,因此每对只需要 12 个字节 - 8 个字节用于 int(在 linux 上,在 Windows 上它甚至可能只有 4 个,但不确定)和 4 个用于 unicode 字符。如果它仍然太大,您可以尝试对其进行分块 - 失败,据我所知,您仍然坚持使用 itertools / generator 方法。
猜你喜欢
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 2012-05-30
  • 2020-05-05
  • 2015-10-28
  • 2016-06-08
  • 1970-01-01
相关资源
最近更新 更多