使用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
相关产品推荐
相关产品推荐

