【问题标题】:Creating fast RGB look up tables in Python在 Python 中创建快速 RGB 查找表
【发布时间】:2018-09-24 20:25:20
【问题描述】:

我有一个函数,我将调用 'rgb2something' 将 RGB 数据 [1x1x3] 转换为单个值(概率),循环遍历输入 RGB 数据中的每个像素结果相当慢。

我尝试了以下方法来加快转换速度。生成 LUT(查找表):

import numpy as np

levels = 256
levels2 = levels**2
lut = [0] * (levels ** 3)

levels_range = range(0, levels)

for r in levels_range:
    for g in levels_range:
        for b in levels_range:
            lut[r + (g * levels) + (b * levels2)] = rgb2something(r, g, b)

并将 RGB 转换为转换后的概率图像:

result = np.take(lut, r_channel + (g_channel * 256) + (b_channel * 65536))

但是,生成 LUT 和计算结果仍然很慢。在 2 维中它相当快,但是在 3 维(r、g 和 b)中它很慢。我怎样才能提高它的性能?

编辑

rgb2something(r, g, b) 看起来像这样:

def rgb2something(r, g, b):
    y = np.array([[r, g, b]])
    y_mean = np.mean(y, axis=0)
    y_centered = y - y_mean
    y_cov = y_centered.T.dot(y_centered) / len(y_centered)
    m = len(Consts.x)
    n = len(y)
    q = m + n
    pool_cov = (m / q * x_cov) + (n / q * y_cov)
    inv_pool_cov = np.linalg.inv(pool_cov)
    g = Consts.x_mean - y_mean
    mah = g.T.dot(inv_pool_cov).dot(g) ** 0.5
    return mah

编辑 2:

我正在尝试实现的完整工作代码示例,我正在使用 OpenCV,因此欢迎使用任何 OpenCV 方法,例如 Apply LUT,以及 C/C++ 方法:

import matplotlib.pyplot as plt
import numpy as np 
import cv2

class Model:
    x = np.array([
        [6, 5, 2],
        [2, 5, 7],
        [6, 3, 1]
    ])
    x_mean = np.mean(x, axis=0)
    x_centered = x - x_mean
    x_covariance = x_centered.T.dot(x_centered) / len(x_centered)
    m = len(x)
    n = 1  # Only ever comparing to a single pixel
    q = m + n
    pooled_covariance = (m / q * x_covariance)  # + (n / q * y_cov) -< Always 0 for a single point
    inverse_pooled_covariance = np.linalg.inv(pooled_covariance)

def rgb2something(r, g, b):
    #Calculates Mahalanobis Distance between pixel and model X
    y = np.array([[r, g, b]])
    y_mean = np.mean(y, axis=0)
    g = Model.x_mean - y_mean
    mah = g.T.dot(Model.inverse_pooled_covariance).dot(g) ** 0.5
    return mah

def generate_lut():
    levels = 256
    levels2 = levels**2
    lut = [0] * (levels ** 3)

    levels_range = range(0, levels)

    for r in levels_range:
        for g in levels_range:
            for b in levels_range:
                lut[r + (g * levels) + (b * levels2)] = rgb2something(r, g, b)

    return lut

def calculate_distance(lut, input_image):
    return np.take(lut, input_image[:, :, 0] + (input_image[:, :, 1] * 256) + (input_image[:, :, 2] * 65536))

lut = generate_lut()
rgb = np.random.randint(255, size=(1080, 1920, 3), dtype=np.uint8)
result = calculate_distance(lut, rgb)

cv2.imshow("Example", rgb)
cv2.imshow("Result", result)
cv2.waitKey(0)

【问题讨论】:

  • 你试过 numpy vectorise 和组合吗?
  • rgb2something 看起来像什么? |是的,生成和使用 64 或 128 MB 的查找表效率不会很高(在 Python 中生成是 2^24 次迭代——解释器很慢——而且巨大的查找对缓存不是很友好)。如上述评论所述,矢量化方法(具有良好、可预测的访问模式)会好得多。
  • 酷,我去看看。是的,我的意思是内存访问模式,内存是线性的。由于 CPU + 缓存的工作方式,最好按顺序(或以小步骤)读取内存 - 它确保大多数时候您需要的数据非常接近 CPU。另一方面,当您以随机模式访问非常大的内存块时,您需要的数据很有可能只在主内存中......这通常需要花费大约 100-150 个时钟周期来获取。
  • @RaymondTunstill 为什么不制作一个字典(本质上是一个哈希表),在需要时为每个 (r, g, b) 元组计算 rgb2something,并将其添加到字典中。这更简单,更易读。此外,在运行时,您最终可能只为一小部分 rgb 值计算 rgb2something 值,而不是像您当前所做的那样为所有 rgb 值计算它们。当然,这取决于您的具体用例。
  • @RohanSaxena 对于一个足够小的图像(少于生成 LUT 的 2^24 次迭代),只运行计算肯定会更快......虽然仍在解释器中运行,所以相对该死的慢。您提到的延迟初始化缓存在这种情况下可能会有所帮助(可能很大程度上取决于输入的类型)。不过,如果我们可以使用矢量化操作而不是在 interpeter 中循环,则可能会有一些加快速度的潜力。

标签: python performance numpy opencv lookup-tables


【解决方案1】:

更新:添加了 blas 优化

有几个直接且非常有效的优化:

(1) 向量化,向量化!对这段代码中的所有内容进行矢量化并不难。见下文。

(2) 使用正确的查找,即花哨的索引,而不是np.take

(3) 使用 Cholesky decomp。使用 blas dtrmm 我们可以利用它的三角形结构

这是代码。只需将其添加到 OP 代码的末尾(在 EDIT 2 下)。除非您非常有耐心,否则您可能还想注释掉 lut = generate_lut()result = calculate_distance(lut, rgb) 行以及对 cv2 的所有引用。我还在x 中添加了一个随机行,以使其协方差矩阵非奇异。

class Full_Model(Model):
    ch = np.linalg.cholesky(Model.inverse_pooled_covariance)
    chx = Model.x_mean@ch

def rgb2something_vectorized(rgb):
    return np.sqrt(np.sum(((rgb - Full_Model.x_mean)@Full_Model.ch)**2,  axis=-1))

from scipy.linalg import blas

def rgb2something_blas(rgb):
    *shp, nchan = rgb.shape
    return np.sqrt(np.einsum('...i,...i', *2*(blas.dtrmm(1, Full_Model.ch.T, rgb.reshape(-1, nchan).T, 0, 0, 0, 0, 0).T - Full_Model.chx,))).reshape(shp)

def generate_lut_vectorized():
    return rgb2something_vectorized(np.transpose(np.indices((256, 256, 256))))

def generate_lut_blas():
    rng = np.arange(256)
    arr = np.empty((256, 256, 256, 3))
    arr[0, ..., 0]  = rng
    arr[0, ..., 1]  = rng[:, None]
    arr[1:, ...] = arr[0]
    arr[..., 2] = rng[:, None, None]
    return rgb2something_blas(arr)

def calculate_distance_vectorized(lut, input_image):
    return lut[input_image[..., 2], input_image[..., 1], input_image[..., 0]]

# test code

def random_check_lut(lut):
    """Because the original lut generator is excruciatingly slow,
    we only compare a random sample, using the original code
    """
    levels = 256
    levels2 = levels**2
    lut = lut.ravel()

    levels_range = range(0, levels)

    for r, g, b in np.random.randint(0, 256, (1000, 3)):
        assert np.isclose(lut[r + (g * levels) + (b * levels2)], rgb2something(r, g, b))

import time
td = []
td.append((time.time(), 'create lut vectorized'))
lutv = generate_lut_vectorized()
td.append((time.time(), 'create lut using blas'))
lutb = generate_lut_blas()
td.append((time.time(), 'lookup using np.take'))
res = calculate_distance(lutv, rgb)
td.append((time.time(), 'process on the fly (no lookup)'))
resotf = rgb2something_vectorized(rgb)
td.append((time.time(), 'process on the fly (blas)'))
resbla = rgb2something_blas(rgb)
td.append((time.time(), 'lookup using fancy indexing'))
resv = calculate_distance_vectorized(lutv, rgb)
td.append((time.time(), None))

print("sanity checks ... ", end='')
assert np.allclose(res, resotf) and np.allclose(res, resv) \
    and np.allclose(res, resbla) and np.allclose(lutv, lutb)
random_check_lut(lutv)
print('all ok\n')

t, d = zip(*td)
for ti, di in zip(np.diff(t), d):
    print(f'{di:32s} {ti:10.3f} seconds')

示例运行:

sanity checks ... all ok

create lut vectorized                 1.116 seconds
create lut using blas                 0.917 seconds
lookup using np.take                  0.398 seconds
process on the fly (no lookup)        0.127 seconds
process on the fly (blas)             0.069 seconds
lookup using fancy indexing           0.064 seconds

我们可以看到,最佳查找胜过最佳即时计算。也就是说,该示例可能高估了查找成本,因为随机像素可能不如自然图像对缓存友好。

原始答案(也许对某些人仍然有用)

如果 rgb2something 无法矢量化,并且您想处理一张典型图像,那么您可以使用 np.unique 获得不错的加速。

如果 rgb2something 很昂贵并且必须处理多个图像,那么 unique 可以与缓存相结合,这可以使用 functools.lru_cache 方便地完成---仅(次要)绊脚石:参数必须是可散列的。事实证明,这种强制的代码修改(将 rgb 数组转换为 3 字节字符串)恰好有益于性能。

只有当您有大量像素覆盖大多数色调时,才值得使用完整的查找表。在这种情况下,最快的方法是使用 numpy 花式索引来进行实际查找。

import numpy as np
import time
import functools

def rgb2something(rgb):
    # waste some time:
    np.exp(0.1*rgb)
    return rgb.mean()

@functools.lru_cache(None)
def rgb2something_lru(rgb):
    rgb = np.frombuffer(rgb, np.uint8)
    # waste some time:
    np.exp(0.1*rgb)
    return rgb.mean()

def apply_to_img(img):
    shp = img.shape
    return np.reshape([rgb2something(x) for x in img.reshape(-1, shp[-1])], shp[:2])

def apply_to_img_lru(img):
    shp = img.shape
    return np.reshape([rgb2something_lru(x) for x in img.ravel().view('S3')], shp[:2])

def apply_to_img_smart(img, print_stats=True):
    shp = img.shape
    unq, bck = np.unique(img.reshape(-1, shp[-1]), return_inverse=True, axis=0)
    if print_stats:
        print('total no pixels', shp[0]*shp[1], '\nno unique pixels', len(unq))
    return np.array([rgb2something(x) for x in unq])[bck].reshape(shp[:2])

def apply_to_img_smarter(img, print_stats=True):
    shp = img.shape
    unq, bck = np.unique(img.ravel().view('S3'), return_inverse=True)
    if print_stats:
        print('total no pixels', shp[0]*shp[1], '\nno unique pixels', len(unq))
    return np.array([rgb2something_lru(x) for x in unq])[bck].reshape(shp[:2])

def make_full_lut():
    x = np.empty((3,), np.uint8)
    return np.reshape([rgb2something(x) for x[0] in range(256)
                       for x[1] in range(256) for x[2] in range(256)],
                      (256, 256, 256))

def make_full_lut_cheat(): # for quicker testing lookup
    i, j, k = np.ogrid[:256, :256, :256]
    return (i + j + k) / 3

def apply_to_img_full_lut(img, lut):
    return lut[(*np.moveaxis(img, 2, 0),)]

from scipy.misc import face

t0 = time.perf_counter()
bw = apply_to_img(face())
t1 = time.perf_counter()
print('naive                 ', t1-t0, 'seconds')

t0 = time.perf_counter()
bw = apply_to_img_lru(face())
t1 = time.perf_counter()
print('lru first time        ', t1-t0, 'seconds')

t0 = time.perf_counter()
bw = apply_to_img_lru(face())
t1 = time.perf_counter()
print('lru second time       ', t1-t0, 'seconds')

t0 = time.perf_counter()
bw = apply_to_img_smart(face(), False)
t1 = time.perf_counter()
print('using unique:         ', t1-t0, 'seconds')

rgb2something_lru.cache_clear()

t0 = time.perf_counter()
bw = apply_to_img_smarter(face(), False)
t1 = time.perf_counter()
print('unique and lru first: ', t1-t0, 'seconds')

t0 = time.perf_counter()
bw = apply_to_img_smarter(face(), False)
t1 = time.perf_counter()
print('unique and lru second:', t1-t0, 'seconds')

t0 = time.perf_counter()
lut = make_full_lut_cheat()
t1 = time.perf_counter()
print('creating full lut:    ', t1-t0, 'seconds')

t0 = time.perf_counter()
bw = apply_to_img_full_lut(face(), lut)
t1 = time.perf_counter()
print('using full lut:       ', t1-t0, 'seconds')

print()
apply_to_img_smart(face())

import Image
Image.fromarray(bw.astype(np.uint8)).save('bw.png')

示例运行:

naive                  6.8886632949870545 seconds
lru first time         1.7458112589956727 seconds
lru second time        0.4085628940083552 seconds
using unique:          2.0951434450107627 seconds
unique and lru first:  2.0168916099937633 seconds
unique and lru second: 0.3118703299842309 seconds
creating full lut:     151.17599205300212 seconds
using full lut:        0.12164952099556103 seconds

total no pixels 786432 
no unique pixels 134105

【讨论】:

  • 感谢您的测试,但是,众所周知,完整的 LUT 更快,问题是如何通过创建较小的 LUT 或为较大的 LUT 创建更快的访问模式来优化此问题。由于这将在实时系统中进行,因此可能会有大范围的色调。
  • @RaymondTunstill “众所周知,完整的 LUT 更快”。好吧,也许对你来说,但你现在才分享了一些关键的信息。无论如何,我怀疑您会发现更新后的答案很有趣。
【解决方案2】:

首先,请在您的rgb2something 函数中添加Consts 是什么,因为这将有助于我们了解该函数的具体作用。

加快此过程的最佳方法是将操作矢量化。

1) 无缓存

您不需要为此操作构建查找表。如果您有一个应用于每个(r, g, b) 向量的函数,您可以使用np.apply_along_axis 简单地将其应用于图像中的每个向量。在以下示例中,我假设 rgb2something 的简单定义作为占位符 - 此函数当然可以替换为您的定义。

def rgb2something(vector):
    return sum(vector)

image = np.random.randint(0, 256, size=(100, 100, 3), dtype=np.uint8)
transform = np.apply_along_axis(rgb2something, -1, image)

这采用image 数组,并将函数rgb2something 应用于沿轴-1(即最后一个通道轴)的每个一维切片。

2) 惰性填充查找表

虽然缓存不是必需的,但在某些特定用例中,它可能会让您受益匪浅。也许您想在数千张图像上执行rgb2something 的这种逐像素操作,并且您怀疑许多像素值会在图像中重复。在这种情况下,构建查找表可以显着提高性能。我建议懒洋洋地填满表格(我建议假设您的数据集跨越的图像有些相似 - 具有相似的对象、纹理等,这意味着它们总共只跨越整个 2 中相对较小的子集^24 搜索空间)。如果您觉得它们跨越了一个相对较大的子集,您可以预先构建整个查找表(请参阅下一节)。

lut = [-1] * (256 ** 3)

def actual_rgb2something(vector):
    return sum(vector)

def rgb2something(vector):
    value = lut[vector[0] + vector[1] * 256 + vector[2] * 65536]

    if value == -1:
        value = actual_rgb2something(vector)
        lut[vector[0] + vector[1] * 256 + vector[2] * 65536] = value

    return value

然后您可以像以前一样转换每个图像:

image = np.random.randint(0, 256, size=(100, 100, 3), dtype=np.uint8)
transform = np.apply_along_axis(rgb2something, -1, image)

3) 预计算缓存

也许您的图像足够多样化,足以涵盖整个搜索范围的一大组,并且整个缓存的构建成本可以通过降低的查找成本来分摊。

from itertools import product

lut = [-1] * (256 ** 3)

def actual_rgb2something(vector):
    return sum(vector)

def fill(vector):
    value = actual_rgb2something(vector)
    lut[vector[0] + vector[1] * 256 + vector[2] * 65536] = value

# Fill the table
total = list(product(range(256), repeat=3))
np.apply_along_axis(fill, arr=total, axis=1)

现在无需再次计算这些值,您只需从表格中查找它们即可:

def rgb2something(vector):
    return lut[vector[0] + vector[1] * 256 + vector[2] * 65536]

转换图片当然和以前一样:

image = np.random.randint(0, 256, size=(100, 100, 3), dtype=np.uint8)
transform = np.apply_along_axis(rgb2something, -1, image)

【讨论】:

  • 我相信操作员希望提高每张图片的性能,而不是每张图片
  • 他的代码也适用于每个可能的 RGB 矢量组合,而不仅仅是沿着图像的 RGB 矢量,即当他沿着他的所有 3 个循环运行时,您已经锁定了沿着轴的操作轴
  • 我添加了另一个带有完整代码示例的编辑,我希望它能够准确地弄清楚 Consts 是什么。只是一个代表模型的静态类成员,并且 rgb2something 将每个像素与该模型进行比较。感谢您的回答,我将看看我使用惰性方法获得了哪些性能提升,但是,用例是每秒大约 40 帧的实时图像处理,因此数据范围会很大。
  • @NiteyaShah 至于你的第一个问题,我只是以数据集中多个图像的用例为例,因为我觉得操作想要在管道中处理多个图像(实际上,根据他对该线程的最后评论,这就是他打算做的事情)。但我仍然建议了不同的技术,它们在不同的用例中可能会更好。
  • @NiteyaShah 对于您的第二条评论:在缓存版本中(特别是在预计算缓存中),我计算每个可能像素的值。请再次查看答案的第(3)部分。
猜你喜欢
  • 2022-11-17
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 2017-07-02
  • 1970-01-01
  • 2011-09-15
相关资源
最近更新 更多