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

如何结合自定义分块函数的并行化与sentence-transformers的encode_multi_process()方法

如何结合自定义分块函数的并行化与sentence-transformers的encode_multi_process()方法

我完全懂你的痛点——先把所有文本都并行分块完再统一调用encode_multi_process(),确实会浪费不少等待时间,尤其是数据集大的时候,分好的块只能躺在内存里闲置,GPU和CPU的资源没法同时跑起来,太可惜了。下面给你几个实用的方案,帮你把分块和编码的流程串起来,最大化利用硬件资源:

方案一:进程池+队列实现流水线式处理

这个方案的核心是让分块和编码两个步骤同时跑起来:开一个进程池专门处理自定义的create_chunks,把分好的块放到线程安全的队列里;同时另一个进程从队列里取批量的块,一旦积累到合适的大小,就立刻交给encode_multi_process()编码。这样分好的块不用等全部分完就能被编码,彻底消除等待时间。

代码示例

首先假设你的自定义分块函数是这样的(你可以替换成自己的逻辑):

def create_chunks(raw_text: str) -> list[str]:
    # 举个简单的分块逻辑:按每50个字符拆分,你可以改成按token数/句子数拆分
    chunk_size = 50
    return [raw_text[i:i+chunk_size] for i in range(0, len(raw_text), chunk_size)]

然后是流水线的实现代码:

from sentence_transformers import SentenceTransformer
import multiprocessing as mp

# 分块工作进程:负责从输入队列取原始文本,分块后放到输出队列
def chunk_worker(input_queue, output_queue):
    while True:
        item = input_queue.get()
        if item is None:  # 收到结束信号就退出
            break
        text_id, raw_text = item
        chunks = create_chunks(raw_text)
        output_queue.put((text_id, chunks))

# 编码工作进程:从输出队列取分好的块,积累到批量后调用encode_multi_process
def encode_worker(output_queue, model, batch_size=32):
    batch = []
    batch_text_ids = []
    while True:
        item = output_queue.get()
        if item is None:  # 处理剩余的最后一批块
            if batch:
                embeddings = model.encode_multi_process(batch, show_progress_bar=False)
                # 这里可以把embeddings和对应的text_id关联,比如存到字典/文件
                print(f"完成一批编码,共{len(batch)}个块")
            break
        text_id, chunks = item
        batch.extend(chunks)
        batch_text_ids.extend([text_id]*len(chunks))
        
        # 积累到设定的批量大小就触发编码
        if len(batch) >= batch_size:
            embeddings = model.encode_multi_process(batch, show_progress_bar=False)
            # 这里根据需求处理编码结果,比如保存到数据库或本地文件
            print(f"完成一批编码,共{len(batch)}个块")
            # 清空批量容器,准备下一批
            batch = []
            batch_text_ids = []

if __name__ == "__main__":
    # 初始化模型
    model = SentenceTransformer('all-MiniLM-L6-v2')
    # 模拟你的数据集:(文本ID, 原始文本)的列表,替换成你自己的数据集
    dataset = [(i, f"这是需要分块编码的示例文本{i},重复多次来模拟长文本。"*10) for i in range(100)]
    
    # 创建进程间安全的队列
    input_queue = mp.Queue(maxsize=10)  # 限制队列大小,避免内存溢出
    output_queue = mp.Queue(maxsize=100)
    
    # 启动分块和编码的工作进程
    chunk_process = mp.Process(target=chunk_worker, args=(input_queue, output_queue))
    encode_process = mp.Process(target=encode_worker, args=(output_queue, model))
    chunk_process.start()
    encode_process.start()
    
    # 把数据集喂入输入队列
    for item in dataset:
        input_queue.put(item)
    
    # 发送结束信号,告诉分块进程所有文本已处理完
    input_queue.put(None)
    chunk_process.join()
    
    # 告诉编码进程所有分块已处理完,等待它处理剩余的块
    output_queue.put(None)
    encode_process.join()

方案二:ProcessPoolExecutor分批处理+即时编码

如果你觉得队列的实现有点复杂,这个方案更简洁:用concurrent.futures.ProcessPoolExecutor并行分块,每处理完N个原始文本,就把对应的分块收集起来调用encode_multi_process(),不用等全部分块完成,同样能让分块和编码的时间部分重叠,提升效率。

代码示例

from sentence_transformers import SentenceTransformer
from concurrent.futures import ProcessPoolExecutor

def create_chunks(raw_text: str) -> list[str]:
    # 替换成你自己的分块逻辑
    chunk_size = 50
    return [raw_text[i:i+chunk_size] for i in range(0, len(raw_text), chunk_size)]

if __name__ == "__main__":
    model = SentenceTransformer('all-MiniLM-L6-v2')
    # 模拟你的数据集,替换成真实数据
    dataset = [f"这是需要分块编码的示例文本{i},重复多次来模拟长文本。"*10 for i in range(100)]
    batch_size = 16  # 每处理16个原始文本就编码一次对应的分块,可根据硬件调整
    all_embeddings = []
    
    with ProcessPoolExecutor() as executor:
        # 并行提交所有分块任务,得到分块结果的迭代器
        chunk_results = executor.map(create_chunks, dataset)
        temp_chunks = []
        for idx, chunks in enumerate(chunk_results):
            temp_chunks.extend(chunks)
            # 每处理完batch_size个原始文本,就编码一次积累的分块
            if (idx + 1) % batch_size == 0:
                print(f"正在编码第{idx//batch_size + 1}批")
                embeddings = model.encode_multi_process(temp_chunks, show_progress_bar=False)
                all_embeddings.extend(embeddings)
                temp_chunks = []
        # 处理剩余的最后一批分块
        if temp_chunks:
            print("正在编码最后一批")
            embeddings = model.encode_multi_process(temp_chunks, show_progress_bar=False)
            all_embeddings.extend(embeddings)

关键注意事项

  • 内存控制:如果你的分块数量极大,要注意不要让临时存储分块的容器(比如方案二中的temp_chunks)积累太大,不然会爆内存。可以根据你的GPU内存大小调整batch_size——GPU内存大就设大一点,内存小就设小一点。
  • 资源竞争规避:encode_multi_process()本身已经会自动利用多CPU和GPU资源,所以编码阶段不要再额外开进程,不然会导致CPU/GPU资源被争抢,反而变慢。
  • 结束信号处理:方案一用队列的方式时,一定要记得发送None作为结束信号,不然工作进程会一直阻塞在get()方法上,无法正常退出。

两种方案各有优劣:方案一更适合超大数据集,流水线式处理能持续利用硬件资源;方案二更简洁易维护,适合中等规模的数据集。你可以根据自己的需求和数据集大小来选择。

备注:内容来源于stack exchange,提问作者Anshu Chen

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.16 07:15:29