Spacy使用en_core_web_trf模型GPU处理文档时出现OOM报错如何解决
问题根因
- 梯度计算中间张量累积:当前代码没有禁用PyTorch的梯度计算,每次推理时产生的反向传播中间张量会一直占用显存,不会自动释放,这是每次调用显存上涨2-3GB的核心原因。
- 长文本推理显存过载:Transformer模型的注意力矩阵显存占用和输入序列长度的平方成正比,直接传入100页长度的整串文本时,单批次输入序列过长,会直接申请数GB的临时显存,触发OOM。
- 版本不兼容问题:spacy 3.1.3要求CUDA版本最低为10.2,你使用的CUDA9.1版本过老,会导致显存管理逻辑异常,进一步加剧显存泄漏问题。
- 旧版本已知bug:spacy 3.1.x分支存在Transformer模型推理时显存未正确回收的已知bug,也会导致显存持续上涨。
解决方案
1. 核心代码修复(解决显存泄漏)
在推理逻辑外增加torch.no_grad()上下文管理器禁用梯度计算,推理结束后手动清空CUDA缓存,修改后的代码如下:
import torch import spacy class SpacyExtractor(): def __init__(self): spacy.require_gpu() # 可选:限制单进程显存占用比例,避免显存占满系统报错 torch.cuda.set_per_process_memory_fraction(0.9) self.model = spacy.load('en_core_web_trf', disable=["tagger", "parser", "attribute_ruler", "lemmatizer"]) def get_named_entities(self, text: str): # 禁用梯度计算,避免中间张量累积 with torch.no_grad(): doc = self.model(text) entities = [] for ent in doc.ents: entities.append((ent.text, ent.label_)) # 手动清空未使用的CUDA缓存 torch.cuda.empty_cache() return entities
2. 长文本分批处理
不要直接传入整份长文本,提前将文本按句子或者固定长度(建议单段不超过400个token,预留冗余空间)拆分,分批调用get_named_entities处理后合并实体结果,避免单批次序列过长导致的显存突增。
3. 依赖版本适配
升级CUDA版本到11.x,同时安装匹配版本的PyTorch、spacy和en_core_web_trf模型,确保各依赖版本兼容,避免因版本不匹配导致的显存管理异常。如果无法升级CUDA,可降级spacy到3.0.x的兼容版本使用。
4. 可选优化配置
如果显存仍然紧张,可以在加载模型时设置Transformer的批次大小参数,降低单批次处理的最大长度,进一步控制显存占用。
内容的提问来源于stack exchange,提问作者Honza
相关产品推荐
相关产品推荐

