Huggingface BertForMaskedLM迭代90次以上无响应无报错问题求助
问题根因及修复方案
根因分析
- M1芯片Mac运行PyTorch默认启用Metal后端,该后端在小批量反复推理场景下存在已知内存泄漏问题,不会主动释放中间张量占用的显存,累计到阈值后就会陷入无报错假死状态
- 原代码未关闭梯度计算,推理过程中会持续累积计算图,占用额外内存,加快泄漏阈值的触发
- 每次迭代生成的
masked_arr、output等张量默认驻留显存,未及时释放
修复步骤
- 所有推理逻辑放到
torch.no_grad()上下文管理器中,关闭梯度计算,避免计算图累积 - 每次迭代结束后调用
torch.mps.empty_cache()清空M1的Metal后端缓存,释放无用显存 - 不需要保留的中间张量可以主动转成Python标量或者调用
del删除,进一步释放空间
修复后代码
import torch import numpy as np from transformers import BertForMaskedLM model = BertForMaskedLM.from_pretrained('bert-base-uncased') # 可选:将模型移到MPS后端加速推理 # model = model.to('mps') imp = {} # 全局关闭梯度计算 with torch.no_grad(): for count, id_to_mask in enumerate(unique_ids): # 生成mask后的输入张量 masked_arr = torch.tensor([[103 if x==id_to_mask else x for x in input_ids]]) # 可选:将输入移到MPS后端加速推理 # masked_arr = masked_arr.to('mps') # 模型推理 output = model(masked_arr) # 计算目标词概率和 idxs = np.where(input_ids==id_to_mask)[0] prob = output.logits[0][idxs][:,id_to_mask] sum_prob = prob.sum().item() # 转为Python标量,不保留张量结构 # 迭代计数打印 print(count) # 结果存储 imp[id_to_mask] = sum_prob # 清空MPS后端缓存 if torch.backends.mps.is_available(): torch.mps.empty_cache() # 主动删除中间张量,释放内存 del masked_arr, output, prob
额外排查方案
如果修改后仍出现卡顿,可做以下验证:
- 升级PyTorch和transformers到最新版本,M1适配相关的bug大多在高版本已修复
- 临时将模型和输入都放在CPU上运行,若CPU运行无卡顿,即可确认是Metal后端适配问题,可排除硬件故障可能
内容的提问来源于stack exchange,提问作者Baymax Lim
相关产品推荐
相关产品推荐

