jax.device_put

目录

jax.device_put#

jax.device_put(x, device=None, *, src=None, donate=False, may_alias=None)[源代码]#

x 传输到 device

参数:
  • x – 一个数组、标量或包含它们的(嵌套)标准 Python 容器。

  • device (None | xc.Device | Sharding | P | Format | Any) – (可选)DeviceSharding,或者标准 Python 容器中的(嵌套)Sharding(必须是 x 的树前缀),表示要将 x 传输到的设备。如果指定,结果将被提交到该设备(或设备集)。

  • src (None | xc.Device | Sharding | P | Format | Any) – (可选)DeviceSharding,或者标准 Python 容器中的(嵌套)Sharding(必须是 x 的树前缀),表示 x 当前所属的设备。

  • donate (bool | Any) – 布尔值或标准 Python 容器中的(嵌套)布尔值(必须是 x 的树前缀)。如果为 True,则 x 可以被覆盖并在调用方中标记为已删除。这是一种尽力而为的机制。JAX 会在可能的情况下进行捐赠(donate),否则则不会。如果已捐赠,输入缓冲区(将来)将始终被删除。

  • may_alias (bool | None | Any) – 布尔值、None 或标准 Python 容器中的(嵌套)布尔值(必须是 x 的树前缀)。如果为 False,x 将被复制。如果为 True,根据运行时的实现,x 可能会被别名化(alias)。

返回:

驻留在 device 上的 x 的副本。

如果 device 参数为 None,则当操作数已经在任何设备上时,该操作表现为恒等函数;否则,它会将数据传输到默认设备,且不进行提交。

此函数始终是异步的,即它会立即返回,而不会阻塞调用它的 Python 线程直到任何传输完成。