Compel 提示词嵌入库显存优化实战:告别VRAM泄漏的5个torch.no_grad技巧
【免费下载链接】compelA prompting enhancement library for transformers-type text embedding systems项目地址: https://gitcode.com/gh_mirrors/co/compel
Compel 是一个面向 transformers 类文本嵌入系统(如 Stable Diffusion 系列)的提示词加权与增强库,支持权重语法(term++)、混合(Blend)和长文本拼接。很多新手在用它批量生成图像时,会遇到显存(VRAM)只涨不跌、最后 OOM 崩溃的问题。本文给出 5 个快速上手的显存优化技巧,帮你彻底解决 VRAM 泄漏问题。
先看现象:为什么显存只涨不跌?
在循环里反复调用 Compel 计算提示词嵌入时,如果 PyTorch 的自动求导(autograd)被激活,每次前向传播都会构建计算图并保存中间激活值。这些张量只要还有引用,显存就不会释放——这就是典型的 "VRAM 泄漏":内存随调用次数线性增长,几十次循环后torch.cuda.OutOfMemoryError就会找上门。
技巧1:认识 Compel 内置的 @torch.no_grad()
好消息是,Compel 的核心调用入口已经做了显存保护。在 src/compel/compel.py 中,__call__方法被@torch.no_grad()装饰器包裹:
@torch.no_grad() def __call__(self, text, return_tokenization=False): ...这意味着只要你走标准用法compel(prompt),嵌入计算过程不会构建计算图,中间张量用完即释放,从源头掐断了 VRAM 泄漏。同理,便捷封装类CompelForSD、CompelForSDXL、CompelForFlux(见 src/compel/convenience_wrappers.py)内部全部经由该入口调用,天然安全。
技巧2:手动调用 build_conditioning_tensor 时,自己套一层 no_grad
注意一个细节:低层方法build_conditioning_tensor()没有@torch.no_grad()装饰(见 src/compel/compel.py)。如果你有自定义流程、直接调用它,请务必自己包上上下文管理器:
with torch.no_grad(): embeds = compel.build_conditioning_tensor(prompt)一句话记忆:谁调用嵌入计算,谁负责关闭梯度。这是新手最常踩的坑。
技巧3:控制 device 与 dtype,把嵌入留在需要的地方
Compel构造时提供两个显存相关的参数(见 src/compel/compel.py 的文档说明):
device:指定创建张量的设备。不指定时跟随text_encoder所在设备。dtype_for_device_getter:默认返回torch.float32。若你的管线以 float16 推理,可传入按设备返回torch.float16的回调,嵌入张量体积直接减半。
经验做法:嵌入算完后.cpu()搬到内存,只在真正送入扩散管线时才.to(device),避免嵌入张量长期占据 GPU。
技巧4:批量提交提示词,复用空串缓存
Compel 支持一次性传入提示词列表(见 src/compel/compel.py 的__call__批量分支),内部会对变长张量做对齐填充后拼接。相比循环单条调用,批量调用减少了重复的填充张量创建和 Python 层开销。
另外,内部用于填充的空串嵌入empty_z是惰性缓存的(见 src/compel/embeddings_provider.py),只计算一次,反复使用无额外显存负担——不需要你手动干预。
技巧5:循环生成后及时释放 + 验证清单
批量出图时,遵循以下操作习惯:
- 及时解引用:生成完一张图后,
del embeds, image,必要时import gc; gc.collect(); - 别在扩散循环里重算嵌入:每步 denoise 时重新调用 Compel 是最常见的泄漏源,嵌入只需算一次;
- 用监控验证:
import torch torch.cuda.empty_cache() print(torch.cuda.memory_allocated() / 1024**3, "GB")循环前后打印一次,数值应基本持平。若持续上涨,优先排查是否误开了梯度(回到技巧2)。
| 检查项 | 正确做法 |
|---|---|
| 调用入口 | 走compel(prompt)或便捷封装类 |
| 手动低层调用 | 包裹torch.no_grad() |
| 嵌入精度 | 与管线一致,推荐 float16 |
| 调用频率 | 每个 prompt 只算一次,跨步复用 |
| 张量生命周期 | 用完del+gc.collect() |
总结
Compel 的显存安全设计(@torch.no_grad()入口、empty_z缓存、无权重时的旁路优化,见 src/compel/embeddings_provider.py)已经帮你挡掉了一半风险;剩下的一半靠使用习惯:手动调用时关梯度、精度与设备按需控制、批量调用、及时释放。做到这 5 点,长时批量出图也能保持显存曲线平稳。更多语法特性可查阅官方文档 doc/README.md。
【免费下载链接】compelA prompting enhancement library for transformers-type text embedding systems项目地址: https://gitcode.com/gh_mirrors/co/compel
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考