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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 02:50:10