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

使用BERT计算Embedding时遭遇计算过载问题求助

问题:BERT计算句子Embedding时循环处理导致机器过载

我尝试用BERT计算句子Embedding,计划用均值池化得到句子向量,但当前代码计算成本极高,循环处理数据时直接导致机器过载,就算用Colab的80GB内存机器也解决不了,求优化方案。


原安装BERT代码

import torch
from transformers import AutoTokenizer, AutoModel
tokenizer = AutoTokenizer.from_pretrained("bert-base-uncased")
model = AutoModel.from_pretrained("bert-base-uncased")

原获取Embedding函数

# 获取BERT词嵌入
def get_word_embedding(text:str):
    input_ids = torch.tensor(tokenizer.encode(text)).unsqueeze(0)  # 批次大小为1
    outputs = model(input_ids)
    last_hidden_states = outputs[1]  
    # 注:此处注释错误,outputs[1]是<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>token的输出,并非均值池化结果
    return last_hidden_states[0]

数据情况

文本最大词数为50,需要计算实体+文本拼接后的Embedding。

原运行代码

entity_desc是我的数据集,以下循环会导致机器过载:

entity_embedding = {}
for i in range(len(entity_desc)):
    entity = entity_desc['entity'][i]
    text = entity_desc['text'][i]
    entity += ' ' + text
    entity_embedding[entity_desc['entity_id'][i]] = get_word_embedding(entity)

优化方案

1. 修正均值池化逻辑

你当前代码用的是outputs[1](<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>token输出),并非均值池化。正确的均值池化需取所有有效token的隐藏层输出做均值:

def get_mean_pooled_embedding(text: str, tokenizer, model, device):
    # 编码文本,自动padding、截断
    inputs = tokenizer(
        text,
        return_tensors="pt",
        padding=True,
        truncation=True,
        max_length=50  # 匹配文本最大词数
    ).to(device)
    
    with torch.no_grad():  # 禁用梯度计算,大幅节省内存
        outputs = model(**inputs)
    
    last_hidden = outputs[0]
    # 用attention mask过滤padding token
    mask = inputs['attention_mask'].unsqueeze(-1).expand(last_hidden.size())
    token_embeddings = last_hidden * mask
    # 计算有效token的均值
    pooled_embedding = torch.sum(token_embeddings, dim=1) / torch.clamp(mask.sum(dim=1), min=1e-9)
    return pooled_embedding.squeeze().cpu().numpy()  # 转numpy减少内存占用

2. 批量处理数据(核心优化)

单条处理效率极低,改成批量处理可大幅降低计算成本:

def batch_process_embeddings(entity_desc, tokenizer, model, device, batch_size=32):
    entity_embedding = {}
    total = len(entity_desc)
    
    for start in range(0, total, batch_size):
        end = min(start + batch_size, total)
        batch_data = entity_desc.iloc[start:end]
        
        # 批量拼接实体和文本
        batch_texts = [f"{row['entity']} {row['text']}" for _, row in batch_data.iterrows()]
        
        # 批量编码
        inputs = tokenizer(
            batch_texts,
            return_tensors="pt",
            padding=True,
            truncation=True,
            max_length=50
        ).to(device)
        
        with torch.no_grad():
            outputs = model(**inputs)
        
        last_hidden = outputs[0]
        mask = inputs['attention_mask'].unsqueeze(-1).expand(last_hidden.size())
        token_embeddings = last_hidden * mask
        pooled_embeddings = torch.sum(token_embeddings, dim=1) / torch.clamp(mask.sum(dim=1), min=1e-9)
        
        # 批量存入结果
        for idx, entity_id in enumerate(batch_data['entity_id']):
            entity_embedding[entity_id] = pooled_embeddings[idx].cpu().numpy()
        
        # 清理临时张量,释放显存
        del inputs, outputs, last_hidden, mask, token_embeddings, pooled_embeddings
        torch.cuda.empty_cache()
    
    return entity_embedding

3. 启用GPU加速

将模型和数据移到GPU上,计算速度提升数倍:

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = model.to(device)

4. 额外内存优化技巧

  • 始终用torch.no_grad()包裹模型推理,避免存储梯度信息
  • 计算后及时删除临时张量,调用torch.cuda.empty_cache()清理显存
  • 用numpy数组存储最终Embedding,比PyTorch张量更节省内存
  • 若数据集极大,分块读取数据,不要一次性加载全部到内存

5. 调用优化后的代码

# 初始化设备
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = model.to(device)

# 批量计算Embedding(可根据内存调整batch_size)
entity_embedding = batch_process_embeddings(entity_desc, tokenizer, model, device, batch_size=64)

内容的提问来源于stack exchange,提问作者edamame

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.11 02:05:23