【问题标题】:Numpy, problem with long arraysNumpy,长数组的问题
【发布时间】:2010-12-14 10:19:09
【问题描述】:

我有两个数组(a 和 b),其中 n 个整数元素在 (0,N) 范围内。

错别字:具有 2^n 个整数的数组,其中最大整数取值 N = 3^n

我想计算 a 和 b 中每个元素组合的总和(所有 i,j 的 sum_ij_ = a_i_ + b_j_)。然后取模N(sum_ij_ = sum_ij_ % N),最后计算不同和的频率。

为了用 numpy 快速做到这一点,没有任何循环,我尝试使用 meshgrid 和 bincount 函数。

A,B = numpy.meshgrid(a,b)
A = A + B
A = A % N
A = numpy.reshape(A,A.size)
result = numpy.bincount(A)

现在,问题是我的输入数组很长。当我使用具有 2^13 个元素的输入时,meshgrid 会给我 MemoryError。我想为具有 2^15-2^20 个元素的数组计算这个。

n 在 15 到 20 的范围内

有没有什么巧妙的技巧可以用 numpy 做到这一点?

我们将不胜感激。

-- 乔恩

【问题讨论】:

  • numpy 真的会那么高效吗?我猜你最好在 C++ 中编写自己的函数并尽可能优化。从听起来像 numpy 无法处理那么大的数组。尽管我必须说,如果您有两个包含 2^15 到 2^20 个元素的数组,那么如果您查看它们所有不同的总和,那么您最终会得到一个包含 2^30 到 2^40 个元素的数组。这是很多..
  • @unutbu: N~3^n @liberalkid: 我猜你是对的。我的 c++ 技能不是很好。

标签: python math numpy


【解决方案1】:

根据 jonalm 的评论进行编辑:

jonalm:N~3^n 不是 n~3^N。 N 是 a 中的最大元素,n 是数量 a中的元素。

n 约为 2^20。如果 N 为 ~ 3^n,则 N 为 ~ 3^(2^20) > 10^(500207)。 科学家估计 (http://www.stormloader.com/ajy/reallife.html) 宇宙中只有大约 10^87 个粒子。因此,计算机没有(天真的)方法可以处理大小为 10^(500207) 的 int。

jonalm:不过我对你定义的 pv() 函数有点好奇。 (一世 不要设法运行它,因为 text.find() 没有定义(猜它在另一个 模块))。这个功能是如何工作的,它的优势是什么?

pv 是我为调试变量值而编写的一个小辅助函数。它像 print() 除了当你说 pv(x) 时,它会同时打印文字变量名称(或表达式字符串)、冒号,然后是变量的值。

如果你放

#!/usr/bin/env python
import traceback
def pv(var):
    (filename,line_number,function_name,text)=traceback.extract_stack()[-2]
    print('%s: %s'%(text[text.find('(')+1:-1],var))
x=1
pv(x)

你应该得到一个脚本

x: 1

与打印相比,使用 pv 的适度优势在于它可以节省您的打字时间。而不是必须 写

print('x: %s'%x)

你可以拍下来

pv(x)

当有多个变量要跟踪时,标记变量会很有帮助。 我只是厌倦了把它全部写出来。

pv 函数通过使用 traceback 模块来查看代码行 用于调用 pv 函数本身。 (参见http://docs.python.org/library/traceback.html#module-traceback)该行代码作为字符串存储在变量文本中。 text.find() 是对常用字符串方法 find() 的调用。例如,如果

text='pv(x)'

然后

text.find('(') == 2               # The index of the '(' in string text
text[text.find('(')+1:-1] == 'x'  # Everything in between the parentheses

我假设 n ~ 3^N 和 n~2**20

这个想法是使用模块 N。这减少了数组的大小。 第二个想法(当 n 很大时很重要)是使用 'object' 类型的 numpy ndarrays,因为如果使用整数 dtype,则存在溢出允许的最大整数大小的风险。

#!/usr/bin/env python
import traceback
import numpy as np

def pv(var):
    (filename,line_number,function_name,text)=traceback.extract_stack()[-2]
    print('%s: %s'%(text[text.find('(')+1:-1],var))

您可以将 n 更改为 2**20,但下面我将展示小 n 会发生什么 所以输出更容易阅读。

n=100
N=int(np.exp(1./3*np.log(n)))
pv(N)
# N: 4

a=np.random.randint(N,size=n)
b=np.random.randint(N,size=n)
pv(a)
pv(b)
# a: [1 0 3 0 1 0 1 2 0 2 1 3 1 0 1 2 2 0 2 3 3 3 1 0 1 1 2 0 1 2 3 1 2 1 0 0 3
#  1 3 2 3 2 1 1 2 2 0 3 0 2 0 0 2 2 1 3 0 2 1 0 2 3 1 0 1 1 0 1 3 0 2 2 0 2
#  0 2 3 0 2 0 1 1 3 2 2 3 2 0 3 1 1 1 1 2 3 3 2 2 3 1]
# b: [1 3 2 1 1 2 1 1 1 3 0 3 0 2 2 3 2 0 1 3 1 0 0 3 3 2 1 1 2 0 1 2 0 3 3 1 0
#  3 3 3 1 1 3 3 3 1 1 0 2 1 0 0 3 0 2 1 0 2 2 0 0 0 1 1 3 1 1 1 2 1 1 3 2 3
#  3 1 2 1 0 0 2 3 1 0 2 1 1 1 1 3 3 0 2 2 3 2 0 1 3 1]

wa 保存 a 中 0、1、2、3 的个数 wb保存b中0、1、2、3的个数

wa=np.bincount(a)
wb=np.bincount(b)
pv(wa)
pv(wb)
# wa: [24 28 28 20]
# wb: [21 34 20 25]
result=np.zeros(N,dtype='object')

将 0 视为令牌或筹码。对于 1,2,3 也是如此。

认为 wa=[24 28 28 20] 表示有一个袋子,里面有 24 个 0 筹码、28 个 1 筹码、28 个 2 筹码、20 个 3 筹码。

你有一个 wa-bag 和一个 wb-bag。当您从每个袋子中抽出一个筹码时,您将它们“添加”在一起并形成一个新筹码。你“修改”答案(模 N)。

想象一下从 wb-bag 中取出 1-chip 并将其与 wa-bag 中的每个芯片相加。

1-chip + 0-chip = 1-chip
1-chip + 1-chip = 2-chip
1-chip + 2-chip = 3-chip
1-chip + 3-chip = 4-chip = 0-chip  (we are mod'ing by N=4)

由于 wb 袋子中有 34 个 1 筹码,当你将它们与 wa=[24 28 28 20] 袋子中的所有筹码相加时,你会得到

34*24 1-chips
34*28 2-chips
34*28 3-chips
34*20 0-chips

由于 34 个 1 筹码,这只是部分计数。您还必须处理其他 wb-bag 中的筹码类型,但这向您展示了以下使用的方法:

for i,count in enumerate(wb):
    partial_count=count*wa
    pv(partial_count)
    shifted_partial_count=np.roll(partial_count,i)
    pv(shifted_partial_count)
    result+=shifted_partial_count
# partial_count: [504 588 588 420]
# shifted_partial_count: [504 588 588 420]
# partial_count: [816 952 952 680]
# shifted_partial_count: [680 816 952 952]
# partial_count: [480 560 560 400]
# shifted_partial_count: [560 400 480 560]
# partial_count: [600 700 700 500]
# shifted_partial_count: [700 700 500 600]

pv(result)    
# result: [2444 2504 2520 2532]

这是最终结果:2444 0s、2504 1s、2520 2s、2532 3s。

# This is a test to make sure the result is correct.
# This uses a very memory intensive method.
# c is too huge when n is large.
if n>1000:
    print('n is too large to run the check')
else:
    c=(a[:]+b[:,np.newaxis])
    c=c.ravel()
    c=c%N
    result2=np.bincount(c)
    pv(result2)
    assert(all(r1==r2 for r1,r2 in zip(result,result2)))
# result2: [2444 2504 2520 2532]

【讨论】:

  • 请注意,c %= N 确实有效(并且可能会使用两倍的内存)。
  • @EOL,是的,c %= N 更好。然而,定义c=(a[:]+b[:,np.newaxis]) 意味着你已经输掉了这场战斗,因为这是一个巨大的二维形状数组 (n,n) 而上述解决方案只使用了几个一维形状数组 (N )。
  • 非常感谢您的回答,我喜欢这种方法。但我认为这对我没有帮助,因为数组 a(和 b)中的所有数字都是不同的(没有提到,我的错)。 bincount(a) 将只包含 1 和 0。N~3^n 不是 n~3^N。 N是a中的最大元素,n是a中的元素数。但是,我对您定义的 pv() 函数有点好奇。 (我没有设法运行它,因为 text.find() 没有定义(猜测它在另一个模块中))。这个功能是如何工作的,它的优势是什么?
  • 亲爱的 Ubuntu。我发现我的符号不一致。我真正的意思是 size(a)=2^n(不是我在第一篇文章中写的 n),max(a)=3^n (=N),n 尽可能高。 a[:]+b[:,np.newaxis] %N 可以做到 n=14,但不能更高。我想要 n~20 => max(a)=3^20
【解决方案2】:

检查你的数学,这是你要求的很多空间:

2^20*2^20 = 2^40 = 1 099 511 627 776

如果您的每个元素只有一个字节,那已经是 1 TB 的内存了。

添加一两个循环。这个问题不适合最大化你的内存和最小化你的计算。

【讨论】:

    【解决方案3】:

    尝试分块。你的 meshgrid 是一个 NxN 矩阵,将其阻止到 10x10 N/10xN/10 并且只计算 100 个 bin,最后将它们相加。这样做只使用了大约 1% 的内存。

    【讨论】:

    • 我想这是要走的路,但是有没有一种聪明的方法可以用 numpy 数组来做到这一点。尽量减少 for 循环的使用。
    • 嘿,有一个块的最佳大小吗?
    • 可能是您可以制作的最大块,并且仍将其安全地塞入 ram 中。
    猜你喜欢
    • 2011-03-04
    • 2019-07-30
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多