基于多进程加速Sentence Embeddings大规模文本计算方案咨询
大规模文本Sentence Embeddings并行加速方案(32核CPU)
一、基础加速:批量推理(核心优化)
现有代码逐句喂入模型,完全浪费了CPU的并行计算能力,先改成批量处理,这是无需多进程就能获得几十倍提速的最有效手段。
import tensorflow as tf from transformers import AutoTokenizer, TFAutoModel tokenizer = AutoTokenizer.from_pretrained('distilbert-base-uncased-finetuned-sst-2-english') model = TFAutoModel.from_pretrained('distilbert-base-uncased-finetuned-sst-2-english') # 批量tokenize,自动补全和截断保证输入长度统一 encoded_input = tokenizer(sentences, padding=True, truncation=True, return_tensors='tf') # 一次性批量推理 outputs = model(**encoded_input) # 提取<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>token的嵌入(也可根据需求用所有token的均值) sentence_embeddings = outputs.last_hidden_state[:, 0, :].numpy()
二、TensorFlow下的多进程进阶优化
如果数据集过大无法一次性加载,用tf.data.Dataset实现并行预处理+分批推理,自动利用CPU多核心:
import tensorflow as tf from transformers import AutoTokenizer, TFAutoModel tokenizer = AutoTokenizer.from_pretrained('distilbert-base-uncased-finetuned-sst-2-english') model = TFAutoModel.from_pretrained('distilbert-base-uncased-finetuned-sst-2-english') # 将句子列表转为tf数据集 dataset = tf.data.Dataset.from_tensor_slices(sentences) # 定义并行预处理函数 def preprocess(text): encoded = tokenizer(text.numpy().decode('utf-8'), padding=True, truncation=True, return_tensors='tf') return encoded['input_ids'], encoded['attention_mask'] # 包装成TensorFlow兼容的函数 def tf_preprocess(text): return tf.py_function(preprocess, [text], [tf.int32, tf.int32]) # 开启多核心预处理,自动适配CPU核心数 dataset = dataset.map(tf_preprocess, num_parallel_calls=tf.data.AUTOTUNE) # 设置合理批次大小(根据内存调整,如64/128) dataset = dataset.batch(64) # 分批推理并收集结果 sentence_embeddings = [] for input_ids, attention_mask in dataset: outputs = model(input_ids=input_ids, attention_mask=attention_mask) embeds = outputs.last_hidden_state[:, 0, :].numpy() sentence_embeddings.extend(embeds)
注意:TensorFlow不要直接用multiprocessing模块开进程,全局计算图会导致进程冲突,tf.data的num_parallel_calls是官方推荐的CPU并行方式。
三、PyTorch版本(更适配CPU多进程)
如果TensorFlow多进程始终有问题,换成PyTorch版本,动态图模型在CPU多进程下兼容性更好,能轻松跑满32核:
方案1:手动分块+进程池
import torch from transformers import AutoTokenizer, AutoModel from multiprocessing import Pool, cpu_count # 每个进程单独加载模型,避免序列化冲突 def init_worker(): global tokenizer, model tokenizer = AutoTokenizer.from_pretrained('distilbert-base-uncased-finetuned-sst-2-english') model = AutoModel.from_pretrained('distilbert-base-uncased-finetuned-sst-2-english') model.eval() # 处理单块句子 def process_chunk(chunk): encoded_input = tokenizer(chunk, padding=True, truncation=True, return_tensors='pt') with torch.no_grad(): outputs = model(**encoded_input) return outputs.last_hidden_state[:, 0, :].numpy() # 将句子列表均分N块(N为CPU核心数) def split_list(lst, n): k, m = divmod(len(lst), n) return [lst[i*k + min(i, m):(i+1)*k + min(i+1, m)] for i in range(n)] if __name__ == '__main__': num_workers = cpu_count() sentence_chunks = split_list(sentences, num_workers) # 启动进程池并行处理 with Pool(num_workers, initializer=init_worker) as pool: results = pool.map(process_chunk, sentence_chunks) # 合并结果 sentence_embeddings = [] for res in results: sentence_embeddings.extend(res)
方案2:用datasets库简化多进程逻辑
from datasets import Dataset from transformers import AutoTokenizer, AutoModel import torch tokenizer = AutoTokenizer.from_pretrained('distilbert-base-uncased-finetuned-sst-2-english') model = AutoModel.from_pretrained('distilbert-base-uncased-finetuned-sst-2-english') model.eval() # 转换为datasets格式 dataset = Dataset.from_dict({"sentences": sentences}) # 定义批量推理函数 def compute_embeddings(batch): encoded = tokenizer(batch["sentences"], padding=True, truncation=True, return_tensors='pt') with torch.no_grad(): outputs = model(**encoded) batch["embeddings"] = outputs.last_hidden_state[:, 0, :].numpy() return batch # 开启32核并行处理,设置合理批次大小 dataset = dataset.map(compute_embeddings, batched=True, batch_size=64, num_proc=32) # 提取最终嵌入结果 sentence_embeddings = dataset["embeddings"]
四、额外提速技巧
- 换用轻量模型:比如
sentence-transformers/all-MiniLM-L6-v2,速度是DistilBERT的2-3倍,嵌入质量接近,专为句子嵌入优化。 - 模型量化:用
torch.quantization对PyTorch模型做INT8量化,可进一步提升CPU推理速度,几乎无精度损失。 - 调整批次大小:根据CPU内存上限调整batch_size(如128/256),最大化利用算力。
内容的提问来源于stack exchange,提问作者akd
相关产品推荐
相关产品推荐

