CLIP模型特征归一化引发GPU显存持续增长的解决方法咨询
CLIP模型计算Embedding时GPU显存持续增长的解决方案
问题描述
我正在使用OpenAI CLIP模型计算clip embeddings,代码如下:
import torch import clip from PIL import Image import urllib.request model, preprocess = clip.load("ViT-B/32", device="cuda") urllib.request.urlretrieve("https://upload.wikimedia.org/wikipedia/commons/3/3a/Cat03.jpg", "cat") def foo(): with Image.open("cat") as img: image = preprocess(img).unsqueeze(0).to("cuda") image_features = model.encode_image(image) image_features /= image_features.norm(dim=-1, keepdim=True) for _ in range(10): foo() print(torch.cuda.memory_allocated())
执行后显存占用持续增长:
380955648 402444800 423933952 445423104 466912256 488401408 509890560 531379712 552868864 574358016
尝试在foo函数末尾添加del image、del image_features和torch.cuda.empty_cache()仅能略微缓解,确认问题由image_features /= image_features.norm(dim=-1, keepdim=True)这行原地操作导致,移除该行后显存不再变化。需要解决大量计算时的显存耗尽问题。
解决方案
1. 替换原地操作为非原地赋值
原地操作(/=)会修改原张量并保留计算图依赖,导致中间资源无法释放。改用非原地赋值让原张量可被正确回收:
def foo(): with Image.open("cat") as img: image = preprocess(img).unsqueeze(0).to("cuda") image_features = model.encode_image(image) # 用非原地赋值替代原地操作 image_features = image_features / image_features.norm(dim=-1, keepdim=True)
2. 用torch.no_grad()包裹推理逻辑
CLIP的编码函数默认会构建计算图,即使推理阶段也会存储梯度相关的中间张量。添加torch.no_grad()可以禁用梯度计算,彻底避免这类显存占用:
def foo(): with Image.open("cat") as img: image = preprocess(img).unsqueeze(0).to("cuda") with torch.no_grad(): image_features = model.encode_image(image) image_features = image_features / image_features.norm(dim=-1, keepdim=True)
3. 显式分离张量切断计算图关联
如果需要保留张量数据但不需要梯度,可以通过detach()将张量从计算图中分离,让相关资源被及时释放:
def foo(): with Image.open("cat") as img: image = preprocess(img).unsqueeze(0).to("cuda") with torch.no_grad(): image_features = model.encode_image(image) image_features = image_features / image_features.norm(dim=-1, keepdim=True) # 分离张量,解除与计算图的绑定 image_features = image_features.detach()
4. 批量处理减少显存分配开销
针对大量图片的场景,批量加载和处理能减少重复的显存分配/释放操作,大幅提升显存利用率:
# 示例:批量处理多张图片 image_paths = ["cat"] * 10 # 构建批量张量 batch_images = torch.stack([preprocess(Image.open(path)).to("cuda") for path in image_paths]) with torch.no_grad(): batch_features = model.encode_image(batch_images) batch_features = batch_features / batch_features.norm(dim=-1, keepdim=True)
5. 升级PyTorch版本
部分旧版本PyTorch在CUDA张量原地操作上存在显存泄漏的底层问题,升级到PyTorch 2.x及以上的稳定版本可修复这类问题。
内容的提问来源于stack exchange,提问作者Luca9984
相关产品推荐
相关产品推荐

