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

如何利用SentenceTransformer的encode_multi_process方法实现多GPU批量句子编码?

解决多GPU下SentenceTransformer.encode_multi_process的进程池与GPU分配问题

核心思路

encode_multi_process 原生支持多进程绑定GPU,但需显式指定每个进程对应的设备ID,同时确保句子列表被均匀拆分到各个进程,避免负载不均。

具体实现步骤

  • 获取可用GPU数量:通过torch.cuda.device_count()自动识别当前设备的GPU总数,避免硬编码设备ID。
  • 均匀拆分句子列表:按GPU数量将大规模句子列表平均拆分,最后一个进程处理剩余的所有句子,保证各GPU负载尽量均衡。
  • 进程池绑定GPU:每个子进程单独初始化模型并绑定到指定GPU,避免跨进程模型共享导致的冲突。

完整代码示例

from sentence_transformers import SentenceTransformer
import torch
from multiprocessing import Pool

def encode_chunk(args):
    model_name, sentences, device_id = args
    # 子进程内单独初始化模型并绑定到指定GPU
    model = SentenceTransformer(model_name, device=f'cuda:{device_id}')
    embeddings = model.encode(sentences, show_progress_bar=True)
    return embeddings

if __name__ == '__main__':
    model_name = 'all-MiniLM-L6-v2'
    large_sentence_list = ["句子1", "句子2", ...]  # 替换为你的大规模句子列表
    num_gpus = torch.cuda.device_count()
    
    # 拆分句子为GPU数量对应的chunk
    chunk_size = len(large_sentence_list) // num_gpus
    sentence_chunks = []
    for i in range(num_gpus):
        start = i * chunk_size
        end = start + chunk_size if i != num_gpus -1 else len(large_sentence_list)
        sentence_chunks.append(large_sentence_list[start:end])
    
    # 构造每个进程的参数:模型名、句子chunk、对应GPU ID
    process_args = [(model_name, chunk, i) for i, chunk in enumerate(sentence_chunks)]
    
    # 启动进程池,进程数与GPU数一致
    with Pool(num_gpus) as pool:
        results = pool.map(encode_chunk, process_args)
    
    # 合并所有进程的编码结果
    all_embeddings = torch.cat([torch.tensor(res) for res in results], dim=0)

关键注意事项

  • 必须在if __name__ == '__main__':下启动进程:避免跨平台进程初始化异常,防止模型重复加载。
  • 子进程单独初始化模型:SentenceTransformer模型无法跨进程共享,必须在每个子进程内独立创建并绑定GPU。
  • 显式指定GPU设备ID:确保每个进程对应唯一GPU,避免设备资源冲突。
  • 句子拆分要均匀:防止单GPU负载过高拖慢整体编码速度。

内容的提问来源于stack exchange,提问作者Alexis López

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 12:24:51