并发

并发性#

JAX 对 Python 并发性的支持有限。

客户端可以从不同的 Python 线程中并发调用 JAX API(例如 jit()grad())。

严禁在多个线程中并发操作 JAX 追踪值(trace values)。换句话说,虽然允许从多个线程调用使用 JAX 追踪的函数(例如 jit()),但不得在传递给 jit() 的函数 f 的实现内部使用线程来操作 JAX 值。如果这样做,最可能的结果是导致 JAX 抛出令人费解的错误。

在多控制器(multi-controller)JAX 中,不同进程必须在给定设备上以相同的顺序应用相同的 JAX 操作。如果您在多控制器 JAX 中使用线程,可以使用 thread_guard() 上下文管理器来检测线程可能导致不同进程以不同顺序调度操作的情况,这会导致非确定性的崩溃。当设置了线程保护(thread guard)时,如果从设置线程保护以外的线程调用 JAX 操作,运行时将会报错。