使用LIME可视化微调后的BERT模型为何出现内存错误?
LIME可视化微调BERT时内存占用过高导致运行终止
我正在使用LIME对微调后的BERT模型进行可视化,但不知为何占用内存过高,被系统终止运行。我的代码如下:
tokenizer = BertTokenizer.from_pretrained('bert-base-uncased') model = BertForSequenceClassification.from_pretrained(f'{BASE_PATH}results/{MODEL}/', num_labels=4) def _proba(texts): encodings = tokenizer(texts, truncation=True, padding=True, max_length=250, return_tensors='pt') pred = model(**encodings) softmax = Softmax(dim = 1) prob = softmax(pred.logits).detach().numpy() return prob explainer = LimeTextExplainer(class_names=['A', 'B', 'C', 'D']) idx = 0 exp = explainer.explain_instance(test_texts[idx], _proba, num_features=4) exp.save_to_file('/lime_vis.html')
我在配备64GB内存的服务器上运行该代码仍出现内存错误,在Colab、Kaggle环境中运行单个示例也会耗尽内存。
问题原因及解决办法
核心原因
LIME的explain_instance默认生成5000个扰动样本,每个样本都要经过BERT模型推理。BERT本身参数量大,加上长文本(max_length=250)的张量维度高,批量处理这些样本时会快速耗尽内存。
具体优化措施
减少扰动样本数量
在explain_instance中设置num_samples参数,降低需要处理的样本量:exp = explainer.explain_instance(test_texts[idx], _proba, num_features=4, num_samples=1000)优化模型推理的内存使用
在_proba函数中禁用梯度计算、及时清理临时张量,避免内存累积:def _proba(texts): encodings = tokenizer(texts, truncation=True, padding=True, max_length=250, return_tensors='pt') # 移到CPU运行(GPU显存不足时优先选择) encodings = {k: v.to('cpu') for k, v in encodings.items()} model.to('cpu') with torch.no_grad(): # 禁用梯度计算,大幅节省内存 pred = model(**encodings) softmax = Softmax(dim=1) prob = softmax(pred.logits).detach().numpy() # 手动清理临时变量 del pred, encodings torch.cuda.empty_cache() # 使用GPU时清理显存 return prob降低文本最大长度
若任务允许,减小max_length以降低单样本张量维度:encodings = tokenizer(texts, truncation=True, padding=True, max_length=128, return_tensors='pt')分批处理扰动样本
将大量扰动样本分成小批次推理,避免一次性加载所有样本:def _proba(texts): batch_size = 32 # 根据内存情况调整 all_probs = [] model.to('cpu') with torch.no_grad(): for i in range(0, len(texts), batch_size): batch_texts = texts[i:i+batch_size] encodings = tokenizer(batch_texts, truncation=True, padding=True, max_length=250, return_tensors='pt') encodings = {k: v.to('cpu') for k, v in encodings.items()} pred = model(**encodings) softmax = Softmax(dim=1) prob = softmax(pred.logits).detach().numpy() all_probs.append(prob) del pred, encodings torch.cuda.empty_cache() return np.concatenate(all_probs, axis=0)
内容的提问来源于stack exchange,提问作者BIPLAB ROY
相关产品推荐
相关产品推荐

