如何限制JAX创建的线程数?配置后线程数无变化
JAX线程数无法通过环境变量控制的问题复现与疑问
问题复现
我尝试复现JAX线程数相关问题,按步骤设置各类环境变量后,JAX创建的线程数仍未改变。
无论执行以下设置了环境变量的代码:
import os, subprocess as sp os.environ["MKL_NUM_THREADS"]="1" os.environ["OPENBLAS_NUM_THREADS"]="1" os.environ["OMP_NUM_THREADS"]="1" os.environ["NUM_INTER_THREADS"]="1" os.environ["NUM_INTRA_THREADS"]="1" os.environ["XLA_FLAGS"]="--xla_cpu_multi_thread_eigen=false intra_op_parallelism_threads=1 --xla_force_host_platform_device_count=1" print(os.sched_getaffinity(0)) import jax print("pre:", int(sp.check_output(f"ls /proc/{os.getpid()}/task | wc -l", shell=True))) jax.numpy.zeros([]) print("post:", int(sp.check_output(f"ls /proc/{os.getpid()}/task | wc -l", shell=True)))
还是执行未设置环境变量的代码:
import os, subprocess as sp # os.environ["MKL_NUM_THREADS"]="1" # os.environ["OPENBLAS_NUM_THREADS"]="1" # os.environ["OMP_NUM_THREADS"]="1" # os.environ["NUM_INTER_THREADS"]="1" # os.environ["NUM_INTRA_THREADS"]="1" # os.environ["XLA_FLAGS"]="--xla_cpu_multi_thread_eigen=false intra_op_parallelism_threads=1 --xla_force_host_platform_device_count=1" # print(os.sched_getaffinity(0)) import jax print("pre:", int(sp.check_output(f"ls /proc/{os.getpid()}/task | wc -l", shell=True))) jax.numpy.zeros([]) print("post:", int(sp.check_output(f"ls /proc/{os.getpid()}/task | wc -l", shell=True)))
输出结果均为:
pre: 1 post: 44
相关疑问
我看到一个相关帖子描述该问题,但不理解其中提出的解决方案:最佳方案应允许用户通过环境变量覆盖主线程池的大小。
内容的提问来源于stack exchange,提问作者desert_ranger
相关产品推荐
相关产品推荐

