You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

无法为JAX CPU并行设置设备数量的问题求助

JAX CPU多核心分片失效:无法识别16线程处理器核心数

我正尝试让JAX程序适配CPU和GPU环境——测试阶段用CPU并行加速,生产阶段用GPU。目前GPU并行可正常运行,但CPU始终无法识别我的16线程(8核心)处理器的核心数,无法实现多核心分片。

我知道XLA_FORCE_HOST_PLATFORM_DEVICE_COUNT需要在JAX初始化前设置,已经在代码最开头尝试了多种设置方式,但都没生效。以下是相关代码片段和Jupyter Notebook输出,求问JAX无法读取该环境变量的原因。

相关代码片段

from multiprocessing import cpu_count
core_count = cpu_count()

### 必须在JAX初始化(包括导入)前设置
# - XLA_FLAGS在jax被导入时读取

# 下面是我试过的几种设置方式

#jax.config.update('xla_force_host_platform_device_count', core_count)
#os.environ["XLA_FORCE_HOST_PLATFORM_DEVICE_COUNT"] = '16'#str(core_count)
#os.environ["XLA_FLAGS"] = '--xla_force_host_platform_device_count=' + str(core_count)
os.environ["XLA_FLAGS"] = f"--xla_force_host_platform_device_count={cpu_count()}"

import jax

# 默认用64位浮点数替代32位,提升精度
jax.config.update('jax_enable_x64', True)
jax.config.update('jax_captured_constants_report_frames', -1)
jax.config.update('jax_captured_constants_warn_bytes', 128 * 1024 ** 2)
jax.config.update('jax_traceback_filtering', 'off')
# https://docs.jax.dev/en/latest/gpu_memory_allocation.html
#jax.config.update('xla_python_client_allocator', '"platform"')
# 没法用jax.config.update设置,改用环境变量
os.environ["XLA_PYTHON_CLIENT_ALLOCATOR"] = '"platform"'

print("\nDefault jax backend:", jax.default_backend())

available_devices = jax.devices()
print(f"Available devices: {available_devices}")

running_device = xla_bridge.get_backend().platform
print("Running device:", running_device, end='')

if running_device == 'cpu':
    print(", with:", core_count, "cores.")

    from jax.sharding import PartitionSpec as P, NamedSharding

    # 创建分片对象来跨设备分配数据:
    # 假设core_count是可用核心设备数
    mesh = jax.make_mesh((core_count,), ('cols',))  # 列维度的1D网格

    # 示例矩阵形状(9, N),比如N=1e7
    #x = jax.random.normal(jax.random.key(0), (9, Np))

    # 指定分片规则:不拆分第0轴(行),跨设备拆分第1轴(列)
    # 然后应用分片规则,将矩阵转换成分片数组,用jax.device_put分配到各设备:
    s0_sharded = jax.device_put(s0, NamedSharding(mesh, P(None, 'cols')))  # 'None'表示不拆分第0轴

    print(s0_sharded.sharding)            # 查看分片规则
    print(s0_sharded.addressable_shards)  # 检查每个设备的分片
    jax.debug.visualize_array_sharding(s0_sharded)

输出结果

Default jax backend: cpu
Available devices: [CpuDevice(id=0)]
Running device: cpu, with: 16 cores.

...

相关代码行:--> 258 mesh = jax.make_mesh((core_count,), ('cols',))  # 1D mesh for columns
... jax后端追踪信息
ValueError: Number of devices 1 must be >= the product of mesh_shape (16,)

问题原因及解决办法

  1. Jupyter内核的状态残留问题
    Jupyter Notebook的内核启动后,若之前运行过导入JAX的代码,或有扩展/预加载模块隐式导入了JAX,会导致你在当前单元格设置的环境变量无法被JAX读取——因为XLA的配置只会在JAX首次初始化时加载。

  2. 环境变量设置的格式与优先级问题

    • 不要把xla_force_host_platform_device_count嵌套在XLA_FLAGS中,直接单独设置XLA_FORCE_HOST_PLATFORM_DEVICE_COUNT环境变量更可靠,避免解析格式错误。
    • jax.config.update('xla_force_host_platform_device_count', ...)的方式必须在导入jax前执行,但在已有JAX导入的内核中无效,因为JAX已经完成初始化。
  3. 正确的修复步骤

    • 完全关闭当前Jupyter内核,重启一个全新的内核。
    • 创建第一个单元格,仅包含环境变量设置和JAX导入,优先运行:
      import os
      from multiprocessing import cpu_count
      os.environ["XLA_FORCE_HOST_PLATFORM_DEVICE_COUNT"] = str(cpu_count())
      import jax
      
    • 运行print(jax.devices())验证,此时应显示16个CpuDevice实例。
    • 之后再运行你的业务代码逻辑。
  4. 脚本运行的额外注意
    若在普通Python脚本中运行,必须确保os.environ的设置语句在任何import jax或导入依赖JAX的模块之前,且脚本运行前没有预加载JAX相关模块。

内容的提问来源于stack exchange,提问作者Kepler7894i

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.12 18:12:42