【问题标题】:Speed up scipy ndimage measurements applied on a numpy 3-d加快应用于 numpy 3-d 的 scipy ndimage 测量
【发布时间】:2017-02-18 22:39:37
【问题描述】:

我有多个大型标记的 numpy 二维数组 (10 000x10 000)。对于每个标签(具有相同编号的连接单元格),我想根据另一个 numpy 3-d 数组(平均值、标准、最大值等)的值计算多个测量值。如果将 3-d numpy 转换为 2-d,则使用 scipy.ndimage.labeled_comprehension 工具可以做到这一点。然而,由于标签的数量和数组的大小相当大,计算需要相当长的时间。我当前的代码似乎是多余的,因为我现在正在为输入图像的每个 3 维迭代相同的标签。我想知道是否有办法加快我的代码速度(例如,将三个 scipy.ndimage.labeled_comprehension 计算合并为一个计算)。

使用形状为 (4200,3000,3) 和 283047 个标签的测试数据集,计算耗时 10:34 分钟


测试数据

example_labels=np.array([[1, 1, 3, 3],
   [1, 2, 2, 3],
   [2, 2, 4, 4],
   [5, 5, 5, 4]])

unique_labels=np.unique(example_labels)
value_array=np.arange(48).reshape(4,4,3)

当前代码和所需输出

def mean_std_measurement(x):
    xmean = x.mean()
    xstd = x.std()
    vals.append([xmean,xstd])

def calculate_measurements(labels, unique_labels, value_arr):
    global vals
    vals=[]
    ndimage.labeled_comprehension(value_array[:,:,0],labels,unique_labels,mean_std_measurement,float,-1)
    val1=np.array(vals)
    vals=[]
    ndimage.labeled_comprehension(value_array[:,:,1],labels,unique_labels,mean_std_measurement,float,-1)
    val2=np.array(vals)
    vals=[]
    ndimage.labeled_comprehension(value_array[:,:,2],labels,unique_labels,mean_std_measurement,float,-1)
    val3=np.array(vals)
    return np.column_stack((unique_labels,val1,val2,val3))

>>> print calculate_measurements(example_labels,unique_labels,value_array)
array([[  1.        ,   5.        ,   5.09901951,   6.        ,
      5.09901951,   7.        ,   5.09901951],
   [  2.        ,  21.        ,   4.74341649,  22.        ,
      4.74341649,  23.        ,   4.74341649],
   [  3.        ,  12.        ,   6.4807407 ,  13.        ,
      6.4807407 ,  14.        ,   6.4807407 ],
   [  4.        ,  36.        ,   6.4807407 ,  37.        ,
      6.4807407 ,  38.        ,   6.4807407 ],
   [  5.        ,  39.        ,   2.44948974,  40.        ,
      2.44948974,  41.        ,   2.44948974]])

【问题讨论】:

  • 这可能影响不大,但您应该在函数开头只计算一次np.unique(labels),并在所有这些函数调用中重用结果。
  • 你说得对,我在脚本中添加了它!
  • 有专门的,更优化的功能scipy.ndimage.minimum, .maximum, .mean, .median, .variance, and .standard_deviation。如果他们提供了您需要的所有统计信息,他们应该会提供健康的加速...如果您需要更多统计信息,可以将图像转换为具有 4 列的 Pandas 数据框并执行groupby(可能会更快,我没有保证)。
  • 使用这些函数,我必须对数组进行六次而不是 3 次的迭代。这使得计算速度变慢。我会尝试研究熊猫的想法。
  • 好吧,这很奇怪。在一台基本的笔记本电脑上,你的代码对我来说只需要 1:31 分钟,而带有 ndimage.mean.standard_deviation 的代码只需要几秒钟。也许您的数据集有些有趣?或者你的时间可能是 10k x 10k 的情况。我用作数据values = np.random.randint(256, size=(4200, 3000, 3))labels = np.random.randint(283047, size=(4200, 3000))。我正在使用 scipy 0.16.1 btw。

标签: python numpy scipy vectorization ndimage


【解决方案1】:

mean_std_measurement 中附加的列表可能会在处理大型数组时形成瓶颈。 labeled_comprehension 将自行返回一个数组,因此第一步就是让 scipy 在后台处理数组构造。

labeled_comprehension 只能应用输出单个值的函数——我怀疑这就是你首先使用列表构造的原因——但我们可以通过让函数输出一个复杂值来欺骗这一点。另一种选择是使用结构化数据类型作为输出,如果我们返回超过 2 个值,这将是必要的。

import numpy as np
from scipy import ndimage
example_labels=np.array([[1, 1, 3, 3],
   [1, 2, 2, 3],
   [2, 2, 4, 4],
   [5, 5, 5, 4]])

unique_labels=np.unique(example_labels)
value_array=np.arange(48).reshape(4,4,3)
# return a complex number to get around labled_comprehension limitations
def mean_std_measurement(x):
    xmean = x.mean()
    xstd = x.std()
    return np.complex(xmean, xstd)

def calculate_measurements(labels, unique_labels, value_array):
    val1 = ndimage.labeled_comprehension(value_array[:,:,0],labels,unique_labels,mean_std_measurement,np.complex,-1)
    val2 = ndimage.labeled_comprehension(value_array[:,:,1],labels,unique_labels,mean_std_measurement,np.complex,-1)
    val3 = ndimage.labeled_comprehension(value_array[:,:,2],labels,unique_labels,mean_std_measurement,np.complex,-1)
    # convert the complex numbers back into reals
    return np.column_stack((np.unique(labels),
                            val1.real,val1.imag,
                            val2.real,val2.imag,
                            val3.real,val3.imag))

【讨论】:

  • “这实际上运行的有点慢……” 你在测量时间时使用了多大的数组? OP 对大小为 (10000, 10000) 的数组感兴趣。大小为 (4, 4) 的数组的时间可能无关紧要。
  • 使用形状为 (4200,3000,3) 和 283047 个标签的测试数据集,计算耗时 10:34 分钟。用你的方法,时间几乎是相同的 10:12 分钟。无论如何,感谢您的努力!
  • 无赖。查看scipy/ndimage/measurements.py 中的源代码,您的原始假设(重复标签三次是瓶颈)似乎是正确的。我建议要么重写函数以处理不同维度的标签和数据数组,要么使用结构化数据类型到处乱搞。 (我尝试了第二个,但没有走这么远)。
猜你喜欢
  • 2015-03-14
  • 2023-03-14
  • 1970-01-01
  • 2015-11-11
  • 1970-01-01
  • 1970-01-01
  • 2015-11-23
  • 2019-09-27
  • 1970-01-01
相关资源
最近更新 更多