Hasty Briefsbeta

双语

JAX backends and devices

14 hours ago
  • JAX默认直接将数据加载到GPU上,当数据集超过VRAM容量时会导致内存不足错误。
  • 除非指定特定后端,否则`jax.devices()`函数仅返回默认后端(如果有GPU则返回GPU,否则返回CPU)的设备。
  • 要检查可用后端,JAX需要尝试不同的后端选项并捕获错误,因为没有内置的列表方法。
  • JAX使用默认设备的概念,可以通过`jax.default_device`上下文管理器临时更改,将数据加载到CPU而非GPU。
  • 解决方案是将数据加载代码包裹在`with jax.default_device(jax.devices("cpu")[0]):`中,以避免GPU内存问题。
  • 加载到CPU的数据可以在训练期间使用`jax.device_put`移动到GPU。