JAX的 JIT编译器会对含有独立语句的列表推导进行并行化吗?

编程语言 2026-07-10

我需要对一组大小各不相同的数组应用一个 jax 函数。比如说:

import jax.numpy as jnp

arrs = [jnp.array([1., 2.]),
        jnp.array([6., 7., 8.]),
        jnp.array([12.])]

fun = jnp.cumsum # Some general function to apply to each section

out = [fun(x) for x in arrs] # Apply function to arrays of different sizes

循环中的每条语句都是相互独立的。jit 编译器会自动尝试并行化这个循环,还是至少以某种方式让循环体并发执行?如果不会,在这种情形下优化执行时间的最佳做法是什么?

附言:就编译时间而言,我不认为有办法在尺寸变化时避免每一步都触发重新编译,但如果有更好的办法也请告诉我。

解决方案

jit上下文之外,列表推导式会被即时执行,每次执行都会异步分派(有关详细信息,请参阅 异步分派)。这意味着Python不必等待设备端的计算完成再继续在循环中分派下一条操作。当操作成本较高且设备端的计算耗时长于Python的分派时间时,这一点就显得很重要,并且后端也能在一定程度上实现并行化。

这里还有一个重要的因素是JIT编译。jnp.cumsum 本身默认是JIT编译的(参见 源码),因此对不同形状或数据类型的每次迭代都会导致重新编译,这会进一步降低代码的分派速度。

在这种情形下,优化执行时间的最佳做法是什么?

要提高这类代码的运行速度,最可取的做法是将数组填充到一个共同形状,将它们堆叠成一个数组,然后使用 jax.vmap 对这些数组进行向量化运算。就你问题中的简单示例而言,这看起来相当可行,但具体取决于实际的运算以及你处理的数组大小分布,可能也并非总是可行。

站内所有文章版权归属LeftHeroAI导航站,无授权禁止任何主体转载、抄袭、复制内容,亦不得私自架设镜像站点。一经侵权,本站将通过法律途径追责。

相关文章