akk*_*kkh 4 python performance jit loops jax
这是一个基本示例。
\n@jax.jit\ndef block(arg1, arg2):\n for x1 in range(cons1):\n for x2 in range(cons2):\n for x3 in range(cons3):\n --do something--\n return result\nRun Code Online (Sandbox Code Playgroud)\n当 cons 很小时,编译时间约为一分钟。缺点较大时,编译时间要长得多\xe2\x80\x9410分钟。我需要更高的缺点。可以做什么?\n根据我所读到的内容,循环是原因。它们在编译时展开。\n有任何解决方法吗?还有jax.fori_loop。但我不明白如何使用它。有 jax.experimental.loops 模块,但我再次无法理解它。
\n我对这一切都很陌生。因此,感谢所有帮助。\n如果您可以提供一些如何使用 jax 循环的示例,我们将不胜感激。
\n另外,什么是好的编译时间?以分钟为单位可以吗?\n在其中一个示例中,编译时间为 262 秒,剩余运行时间约为 0.1-0.2 秒。
\n运行时的任何增益都会被编译时间所掩盖。
\nJAX 的 JIT 编译器扁平化所有 Python 循环。要明白我的意思,请看一下这个简单的函数 run through jax.make_jaxpr,这是一种检查 JAX 的跟踪器如何解释 python 代码的方法(有关更多信息,请参阅了解 Jaxprs ):
import jax
def f(x):
for i in range(5):
x += i
return x
print(jax.make_jaxpr(f)(0))
# { lambda ; a.
# let b = add a 0
# c = add b 1
# d = add c 2
# e = add d 3
# f = add e 4
# in (f,) }
Run Code Online (Sandbox Code Playgroud)
请注意,循环已扁平化:每个步骤都变成发送到 XLA 编译器的显式操作。XLA 编译时间随着函数中操作数量的增加而增加,因此三重嵌套的 for 循环会导致较长的编译时间是有道理的。
那么,如何解决这个问题呢?好吧,不幸的是,答案取决于你--do something--在做什么,所以我无法猜测。
一般来说,最好的选择是使用向量化数组运算,而不是循环这些向量中的值;例如,以下是添加两个向量的非常慢的方法:
import jax.numpy as jnp
def f_slow(x, y):
z = []
for xi, yi in zip(xi, yi):
z.append(xi + yi)
return jnp.array(z)
Run Code Online (Sandbox Code Playgroud)
这是做同样事情的更快的方法:
import jax.numpy as jnp
def f_slow(x, y):
z = []
for xi, yi in zip(xi, yi):
z.append(xi + yi)
return jnp.array(z)
Run Code Online (Sandbox Code Playgroud)
如果您的操作不适合矢量化,另一种选择是使用宽松的控制流运算符来代替for循环:这会将循环向下推入 XLA。这在 CPU 上具有相当好的性能,但与等效的向量化数组操作相比,在加速器上速度较慢。
有关 JAX 和 Python 控制流语句(例如for、if、while等)的更多讨论,请参阅JAX - The Sharp Bits : Control Flow。
| 归档时间: |
|
| 查看次数: |
5560 次 |
| 最近记录: |