使用Hugging Face Transformers获取嵌入时内存不足的解决方案咨询
RoBERTa处理大段落内存不足的优化方案
问题描述
使用RoBERTa-base模型生成段落嵌入以计算相似度时,一次性处理1324行段落,即便拥有25GB内存仍因内存不足失败,代码如下:
from transformers import RobertaTokenizer, RobertaModel import torch device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') tokenizer = RobertaTokenizer.from_pretrained("roberta-base") model = RobertaModel.from_pretrained("roberta-base").to(device) inputs = tokenizer(dict_anrika['Anrika'], return_tensors="pt", truncation=True, padding=True).to(device) outputs = model(**inputs)
优化方法与错误排查
核心错误点
- 未启用推理模式:代码未设置
model.eval(),也未使用torch.no_grad(),训练模式下会保留梯度信息,额外占用大量内存。 - 一次性全量输入:将1324条段落全部输入模型,生成超大张量,直接耗尽内存。
优化方案
批量处理数据
将段落分成小批次(如32/64条一批)循环处理,避免单张张量占用过多内存。示例代码:from transformers import RobertaTokenizer, RobertaModel import torch device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') tokenizer = RobertaTokenizer.from_pretrained("roberta-base") model = RobertaModel.from_pretrained("roberta-base").to(device) model.eval() # 切换到推理模式 paragraphs = dict_anrika['Anrika'] batch_size = 32 # 根据内存情况调整 all_embeddings = [] with torch.no_grad(): # 禁用梯度计算,节省内存 for i in range(0, len(paragraphs), batch_size): batch = paragraphs[i:i+batch_size] inputs = tokenizer(batch, return_tensors="pt", truncation=True, padding=True).to(device) outputs = model(**inputs) # 取<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]> token的嵌入作为段落表示,也可选择取序列均值 embeddings = outputs.last_hidden_state[:, 0, :].cpu() # 移至CPU释放GPU内存 all_embeddings.append(embeddings) final_embeddings = torch.cat(all_embeddings, dim=0)启用推理模式与梯度禁用
必须添加model.eval()和torch.no_grad(),关闭训练相关的梯度计算和参数更新,减少内存消耗。限制序列最大长度
若段落语义不受影响,可将max_length设为更小值(如256),进一步缩小张量规模:inputs = tokenizer(batch, return_tensors="pt", truncation=True, padding=True, max_length=256).to(device)使用半精度推理
支持FP16的GPU可将模型和输入转为float16,内存占用直接减半:model = RobertaModel.from_pretrained("roberta-base").to(device).half() inputs = tokenizer(batch, return_tensors="pt", ...).to(device).half()手动清理内存
每批处理后删除无用张量并清理GPU缓存:with torch.no_grad(): for i in range(0, len(paragraphs), batch_size): batch = paragraphs[i:i+batch_size] inputs = tokenizer(batch, return_tensors="pt", truncation=True, padding=True).to(device) outputs = model(**inputs) embeddings = outputs.last_hidden_state[:, 0, :].cpu() all_embeddings.append(embeddings) # 清理当前批次的张量 del inputs, outputs torch.cuda.empty_cache()检查异常长段落
统计所有段落的token长度,若存在极端长的段落(如单条含数万字符),可将其分段后取均值嵌入,避免单条数据占用过多内存。
内容的提问来源于stack exchange,提问作者akshit bhatia
相关产品推荐
相关产品推荐

