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

Huggingface BertForMaskedLM迭代90次以上无响应无报错问题求助

问题根因及修复方案

根因分析

  • M1芯片Mac运行PyTorch默认启用Metal后端,该后端在小批量反复推理场景下存在已知内存泄漏问题,不会主动释放中间张量占用的显存,累计到阈值后就会陷入无报错假死状态
  • 原代码未关闭梯度计算,推理过程中会持续累积计算图,占用额外内存,加快泄漏阈值的触发
  • 每次迭代生成的masked_arr、output等张量默认驻留显存,未及时释放

修复步骤

  1. 所有推理逻辑放到torch.no_grad()上下文管理器中,关闭梯度计算,避免计算图累积
  2. 每次迭代结束后调用torch.mps.empty_cache()清空M1的Metal后端缓存,释放无用显存
  3. 不需要保留的中间张量可以主动转成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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 12:30:05