JAX中映射函数数组的最高效惯用实现方式问询
JAX中映射函数数组的高效惯用实现及性能分析
背景
有开发者提出了一种用lax.switch结合vmap批量应用多个函数的实现方式,示例代码如下:
from jax import lax, vmap import jax.numpy as jnp def func1(x): return 2 * x def func2(x): return -2 * x def func3(x): return 0 * x functions = [func1, func2, func3] index = jnp.arange(len(functions)) x = jnp.ones((3, 5)) vmap_functions = vmap(lambda i, x: lax.switch(i, functions, x)) vmap_functions(index, x) # DeviceArray([[ 2., 2., 2., 2., 2.], # [-2., -2., -2., -2., -2.], # [ 0., 0., 0., 0., 0.]], dtype=float32)
问题解答
1. 当前方式是否为JAX中映射函数数组的惯用实现?
这种vmap搭配lax.switch的写法是JAX里批量处理函数数组比较惯用且高效的方案,特别适合函数数量固定、需要给批量输入逐一匹配对应函数的场景。
如果函数逻辑可以统一向量化表达,可能还有更简洁的写法,但对于逻辑差异较大、无法合并的函数集合,这种方式是业内普遍认可的实现思路。
2. 该方法的性能损耗分析
编译时损耗
lax.switch会把所有传入的函数都编译进计算图,函数数量越多,编译后的计算图体积就越大,单次JIT编译的耗时会增加,同时也会占用更多内存存储编译后的内核。如果函数数量极多,编译阶段的资源开销会比较明显。
运行时损耗
vmap会自动对批量输入做并行化处理,但lax.switch本质是在每个分支选择对应函数执行。当函数数量较多时,硬件的指令级并行可能会受到一定限制;不过JAX的XLA编译器会尽可能做优化,比如合并相似的函数分支逻辑,所以实际运行时的性能损耗通常不会特别显著,除非函数数量极大或者单个函数的逻辑差异悬殊。
内容的提问来源于stack exchange,提问作者Inquisitive
相关产品推荐
相关产品推荐

