NEP 11 — 延迟 UFunc 计算#
- 作者:
Mark Wiebe <mwwiebe@gmail.com>
- 内容类型:
text/x-rst
- 创建时间:
2010年11月30日
- 状态:
Deferred
摘要#
本 NEP 描述了一项为 NumPy 的 UFunc 添加延迟计算的提案。这将允许像 “a[:] = b + c + d + e” 这样的 Python 表达式在一次遍历所有变量的过程中完成计算,而无需创建临时数组。由此产生的性能可能与 numexpr 库相当,但语法更为自然。
这一想法与 UFunc 错误处理和 UPDATEIFCOPY 标志存在交互,影响了设计与实现,但最终结果允许 Python 用户以极小的额外工作量使用延迟计算。
动机#
NumPy 的 UFunc 执行方式会导致大型表达式的性能不佳,因为在此过程中会分配多个临时数组,并对输入进行多次遍历。numexpr 库通过在适合缓存的小块中执行计算,并按元素计算整个表达式,可以在处理此类大型表达式时超越 NumPy。这使得每个输入只需遍历一次,这对缓存而言显著更优。
关于如何在不更改 Python 代码的情况下在 NumPy 中实现这种行为,可以参考 C++ 的表达式模板(expression templates)技术。这些技术可以用来相当随意地重组使用向量或其他数据结构的表达式,例如
A = B + C + D;
可以转换为等效于
for(i = 0; i < A.size; ++i) {
A[i] = B[i] + C[i] + D[i];
}
这是通过返回一个知道如何计算结果的代理对象,而不是返回实际对象来实现的。利用现代 C++ 优化编译器,所产生的机器码通常与手写循环相同。关于此示例,请参见 Blitz++ 库。最近创建的一个用于辅助编写表达式模板的库是 Boost Proto。
通过在 Python 中使用返回代理对象的相同思想,我们可以动态地实现同样的效果。返回的对象是一个尚未分配缓冲区的 ndarray,并且具备在需要时自行计算的足够信息。当一个“延迟数组”最终被求值时,我们可以使用由所有操作数延迟数组组成的表达式树,从而有效地创建一个即时计算的单一新 UFunc。
Python 代码示例#
以下是如何在 NumPy 中使用它的示例。
# a, b, c are large ndarrays
with np.deferredstate(True):
d = a + b + c
# Now d is a 'deferred array,' a, b, and c are marked READONLY
# similar to the existing UPDATEIFCOPY mechanism.
print d
# Since the value of d was required, it is evaluated so d becomes
# a regular ndarray and gets printed.
d[:] = a*b*c
# Here, the automatically combined "ufunc" that computes
# a*b*c effectively gets an out= parameter, so no temporary
# arrays are needed whatsoever.
e = a+b+c*d
# Now e is a 'deferred array,' a, b, c, and d are marked READONLY
d[:] = a
# d was marked readonly, but the assignment could see that
# this was due to it being a deferred expression operand.
# This triggered the deferred evaluation so it could assign
# the value of a to d.
不过,可能会出现一些令人惊讶的行为。
with np.deferredstate(True):
d = a + b + c
# d is deferred
e[:] = d
f[:] = d
g[:] = d
# d is still deferred, and its deferred expression
# was evaluated three times, once for each assignment.
# This could be detected, with d being converted to
# a regular ndarray the second time it is evaluated.
我认为文档中应推荐的使用方式是:将延迟状态保持为默认值,仅在评估可能从中受益的大型表达式时才开启。
# calculations
with np.deferredstate(True):
x = <big expression>
# more calculations
这将避免因始终保持开启延迟使用而导致的意外,例如在稍后使用延迟表达式时出现浮点警告或异常。通过推荐这种方法,有望避免诸如“为什么我的 print 语句抛出除以零错误?”之类的用户提问。
提议的延迟计算 API#
为了使延迟计算生效,C API 需要意识到其存在,并能够在必要时触发计算。ndarray 将获得两个新标志。
NPY_ISDEFERRED表明该 ndarray 实例的表达式计算已被延迟。
NPY_DEFERRED_WASWRITEABLE仅当
PyArray_GetDeferredUsageCount(arr) > 0时可设置。它表明当arr首次在延迟表达式中使用时,它是一个可写数组。如果设置了此标志,调用PyArray_CalculateAllDeferred()将使arr再次变为可写。
注意
问题
NPY_DEFERRED 和 NPY_DEFERRED_WASWRITEABLE 是否应该对 Python 可见?或者从 Python 访问这些标志时,是否应该在必要时触发 PyArray_CalculateAllDeferred?
该 API 将扩展若干函数。
int PyArray_CalculateAllDeferred()
此函数强制发生所有当前已延迟的计算。
例如,如果错误状态被设置为忽略所有,并且执行了 np.seterr({all=’raise’}),这将改变已经延迟的表达式的行为。因此,在更改错误状态之前,应计算所有现有的延迟数组。
int PyArray_CalculateDeferred(PyArrayObject* arr)
如果 ‘arr’ 是一个延迟数组,则为其分配内存并计算延迟表达式。如果 ‘arr’ 不是延迟数组,则直接返回成功。返回 NPY_SUCCESS 或 NPY_FAILURE。
int PyArray_CalculateDeferredAssignment(PyArrayObject* arr, PyArrayObject* out)
如果 ‘arr’ 是一个延迟数组,则将延迟表达式计算到 ‘out’ 中,且 ‘arr’ 保持为延迟数组。如果 ‘arr’ 不是延迟数组,则将其值复制到 out 中。返回 NPY_SUCCESS 或 NPY_FAILURE。
int PyArray_GetDeferredUsageCount(PyArrayObject* arr)
返回有多少个延迟表达式将此数组作为操作数使用的计数。
Python API 将进行如下扩展。
numpy.setdeferred(state)启用或禁用延迟计算。True 表示始终使用延迟计算。False 表示从不使用延迟计算。None 表示当错误处理状态设置为忽略一切时使用延迟计算。在 NumPy 初始化时,延迟状态为 None。
返回先前的延迟状态。
numpy.getdeferred()
返回当前的延迟状态。
numpy.deferredstate(state)
一种用于处理延迟状态的上下文管理器,类似于
numpy.errstate。
错误处理#
错误处理是延迟计算中一个棘手的问题。如果 NumPy 错误状态为 {all=’ignore’},将延迟计算作为默认值可能是合理的;然而,如果一个 UFunc 能够引发错误,那么由后来的 ‘print’ 语句抛出异常而不是引发错误的实际操作抛出异常,将是非常奇怪的。
一种好的方法可能是默认仅在错误状态设置为忽略所有时启用延迟计算,但允许用户通过 ‘setdeferred’ 和 ‘getdeferred’ 函数进行控制。True 表示始终使用延迟计算,False 表示从不使用,None 表示仅在安全时(即错误状态设置为忽略所有)使用。
与 UPDATEIFCOPY 的交互#
NPY_UPDATEIFCOPY 文档指出
数据区域表示一个(行为良好的)副本,当删除此数组时,其信息应传回原始数组。
这是一个特殊标志,如果此数组代表一个副本(因为用户在 PyArray_FromAny 中需要某些标志,而不得不对另一个数组进行复制,并且用户要求在这种情况下设置此标志),则会设置该标志。base 属性随后指向“行为不端”的数组(该数组被设置为 read_only)。当设置了此标志的数组被销毁时,它会将内容拷回“行为不端”的数组(必要时进行类型转换),并将“行为不端”的数组重置为 NPY_WRITEABLE。如果“行为不端”的数组最初不是 NPY_WRITEABLE,那么 PyArray_FromAny 将会返回错误,因为 NPY_UPDATEIFCOPY 将无法实现。
当前 UPDATEIFCOPY 的实现假设它是唯一以这种方式操作 writeable 标志的机制。这些机制必须相互配合才能正常工作。以下是它们可能出错的示例
对 ‘arr’ 进行带 UPDATEIFCOPY 的临时复制(‘arr’ 变为只读)
在延迟表达式中使用 ‘arr’(延迟使用计数变为一,NPY_DEFERRED_WASWRITEABLE 未设置,因为 ‘arr’ 是只读的)
销毁临时副本,导致 ‘arr’ 变为可写
写入 ‘arr’ 破坏了延迟表达式的值
为了处理这个问题,我们将这两个状态设为互斥。
使用 UPDATEIFCOPY 时会检查
NPY_DEFERRED_WASWRITEABLE标志,如果已设置,则调用PyArray_CalculateAllDeferred以在继续操作前刷新所有延迟计算。ndarray 获得了一个新标志
NPY_UPDATEIFCOPY_TARGET,表明该数组将在未来的某个时刻被更新并变为可写。如果延迟计算机制在任何操作数中看到此标志,它会触发立即计算。
其他实现细节#
当创建一个延迟数组时,它会获得 UFunc 所有操作数的引用,以及 UFunc 本身的引用。每个操作数的 ‘DeferredUsageCount’ 会递增,并在延迟表达式计算完成或延迟数组被销毁后递减。
全局维护一个按创建顺序排列的所有延迟数组的弱引用列表。当调用 PyArray_CalculateAllDeferred 时,最新创建的延迟数组将首先被计算。这可能会释放延迟表达式树中包含的其他延迟数组的引用,这些数组随后将无需再进行计算。
进一步优化#
与其在任何错误未设置为 ‘ignore’ 时保守地禁用延迟计算,不如让每个 UFunc 提供一组它可能生成的错误。那么,如果所有这些错误都设置为 ‘ignore’,即使其他错误未设置为忽略,也可以使用延迟计算。
一旦表达式树被明确存储,就可以对其进行转换。例如,add(add(a,b),c) 可以转换为 add3(a,b,c),或者 add(multiply(a,b),c) 在可用时可以使用 CPU 的乘加融合(fma)指令变为 fma(a,b,c)。
虽然我将延迟计算仅限于 UFunc,但它也可以扩展到其他函数,例如 dot()。例如,链式矩阵乘法可以重新排序以最小化中间结果的大小,或者窥孔优化(peep-hole style)遍历可以搜索匹配优化的 BLAS 或其他高性能库调用的模式。
对于超大数组上的操作,将类似 LLVM 的 JIT 集成到此系统中可能大有裨益。UFuncs 和其他操作将提供位码(bitcode),这些位码可以被内联在一起,由 LLVM 优化器优化,然后执行。事实上,迭代器本身也可以用位码表示,从而允许 LLVM 在进行优化时考虑整个迭代过程。