无法为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,)
问题原因及解决办法
Jupyter内核的状态残留问题
Jupyter Notebook的内核启动后,若之前运行过导入JAX的代码,或有扩展/预加载模块隐式导入了JAX,会导致你在当前单元格设置的环境变量无法被JAX读取——因为XLA的配置只会在JAX首次初始化时加载。环境变量设置的格式与优先级问题
- 不要把
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已经完成初始化。
- 不要把
正确的修复步骤
- 完全关闭当前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实例。 - 之后再运行你的业务代码逻辑。
脚本运行的额外注意
若在普通Python脚本中运行,必须确保os.environ的设置语句在任何import jax或导入依赖JAX的模块之前,且脚本运行前没有预加载JAX相关模块。
内容的提问来源于stack exchange,提问作者Kepler7894i
相关产品推荐
相关产品推荐

