使用BERT生成词嵌入时遇内存不足问题,求解决方案
解决BERT前向传播CUDA内存不足的方案
核心问题分析
你当前一次性将2370条数据全部输入模型,即便max_length设为128,整体张量的显存占用仍超出GPU剩余空间。以下是针对性解决方案:
1. 分批次处理数据
这是最直接有效的方法,将数据集拆分成小批量循环处理,每次仅加载部分数据到显存,处理完成后释放内存再处理下一批。
示例代码:
from torch.utils.data import TensorDataset, DataLoader # 将tokenized数据转为Dataset dataset = TensorDataset(tokenized_texts['input_ids'], tokenized_texts['attention_mask']) # 根据显存调整批量大小(如32、16甚至8) dataloader = DataLoader(dataset, batch_size=32) all_word_embeddings = [] with torch.no_grad(): for batch_input_ids, batch_attn_mask in dataloader: # 移至GPU处理 batch_input_ids = batch_input_ids.to('cuda') batch_attn_mask = batch_attn_mask.to('cuda') outputs = bert_model(batch_input_ids, attention_mask=batch_attn_mask) batch_embeddings = outputs.last_hidden_state # 结果移回CPU避免显存占用 all_word_embeddings.append(batch_embeddings.cpu()) # 合并所有批次结果 word_embeddings = torch.cat(all_word_embeddings, dim=0)
2. 启用梯度检查点
开启BERT的梯度检查点功能,以少量计算时间为代价大幅降低显存占用,适配推理场景。
在模型初始化后添加代码:
bert_model.gradient_checkpointing_enable()
3. 使用半精度(FP16)推理
利用PyTorch自动混合精度,将模型参数和计算转为半精度,显存占用可减少约一半。
示例代码:
from torch.cuda.amp import autocast bert_model = bert_model.to('cuda') all_word_embeddings = [] with torch.no_grad(), autocast(): for batch_input_ids, batch_attn_mask in dataloader: batch_input_ids = batch_input_ids.to('cuda') batch_attn_mask = batch_attn_mask.to('cuda') outputs = bert_model(batch_input_ids, attention_mask=batch_attn_mask) batch_embeddings = outputs.last_hidden_state all_word_embeddings.append(batch_embeddings.cpu()) word_embeddings = torch.cat(all_word_embeddings, dim=0)
4. 优化显存分配与清理
- 设置环境变量减少显存碎片,在代码最开头添加:
import os os.environ['PYTORCH_CUDA_ALLOC_CONF'] = 'max_split_size_mb:128'
- 批次处理后手动清理无用显存:
torch.cuda.empty_cache()
5. 改用轻量版BERT模型
替换为参数量更少的DistilBERT,性能接近BERT但显存占用仅为原模型的60%左右:
from transformers import DistilBertTokenizer, DistilBertModel # 替换分词器和模型 bert_tokenizer = DistilBertTokenizer.from_pretrained('distilbert-base-uncased') bert_model = DistilBertModel.from_pretrained('distilbert-base-uncased')
内容的提问来源于stack exchange,提问作者Lol
相关产品推荐
相关产品推荐

