使用torch.multiprocessing调用HF Transformers时CPU全核占用问题
问题排查与解决
核心原因
使用spawn启动子进程时,主进程的torch.set_num_threads(1)不会被继承,每个子进程会默认启用所有CPU线程;同时numpy依赖的BLAS类库(如OpenBLAS、MKL)默认也会使用多线程,两者叠加导致即便限制了进程数,仍会占满所有CPU核心。
具体解决方案
1. 在子进程中统一限制线程数
由于spawn模式下子进程是全新的Python环境,必须在每个子进程启动时重新配置线程参数,可通过Pool的initializer参数批量设置:
import torch import torch.multiprocessing as mp import os def init_worker(): # 限制PyTorch运算线程数 torch.set_num_threads(1) torch.set_num_interop_threads(1) # 限制numpy/BLAS相关库的线程数 os.environ["OMP_NUM_THREADS"] = "1" os.environ["MKL_NUM_THREADS"] = "1" os.environ["OPENBLAS_NUM_THREADS"] = "1" # 主进程设置启动方式 mp.set_start_method("spawn", force=True) seqs = ["your_seq1", "your_seq2"] # 替换为实际字符串列表 def func1(x): # 你的numpy计算逻辑 import numpy as np a = np.random.rand(10) b = np.random.rand(10) return a, b def func2(a, b, x): # 执行torch计算与模型推理 from your_module import ESM2 # 替换为ESM2类的实际导入路径 inputs = worker_model.tokenizer(x, return_tensors="pt") results = worker_model(inputs) # 额外torch计算逻辑 c = results.logits.detach().numpy() return c # 子进程全局存储模型实例,避免重复加载 worker_model = None def main_func(x): a, b = func1(x) c = func2(a, b, x) return c # 创建进程池时绑定初始化函数 pool = mp.Pool(processes=10, initializer=init_worker) results = pool.map(main_func, seqs) pool.close() pool.join()
2. 优化模型加载逻辑
每个子进程只初始化一次模型,避免重复加载带来的资源浪费和线程冲突:
def init_worker(): global worker_model # 先设置线程限制 torch.set_num_threads(1) torch.set_num_interop_threads(1) os.environ["OMP_NUM_THREADS"] = "1" os.environ["MKL_NUM_THREADS"] = "1" os.environ["OPENBLAS_NUM_THREADS"] = "1" # 子进程启动时加载模型 from your_module import ESM2 worker_model = ESM2()
3. 验证线程配置生效
可以在初始化函数中添加打印,确认子进程的线程设置:
def init_worker(): torch.set_num_threads(1) print(f"子进程 {os.getpid()} 的torch线程数: {torch.get_num_threads()}") # 其他配置...
关键注意点
spawn模式下子进程完全独立,主进程的全局配置不会继承,必须在子进程内重新设置线程参数。- 必须同时限制PyTorch和numpy相关库的线程数,才能彻底解决CPU占满的问题。
- 每个子进程只加载一次模型,既能节省内存,也能避免重复初始化导致的线程异常。
内容的提问来源于stack exchange,提问作者JonnyJack
相关产品推荐
相关产品推荐

