如何清除DiffusionPipeline中UNet2DConditionModel的梯度积累
解决Diffusers 0.3.0中LDM Pipeline的梯度清除问题
问题背景
基于prompt-to-prompt代码,使用CompVis/ldm-text2im-large-256模型,通过以下方式加载Pipeline:
model = DiffusionPipeline.from_pretrained(model_id, height=IMAGE_RES, width=IMAGE_RES).to(device)
在不使用torch.no_grad()的情况下调用text2image_ldm方法时,梯度会在noise_pred = model.unet(latents_input, t, encoder_hidden_states=context)["sample"]行积累。尝试过model.unet.zero_grad()和手动遍历参数设置param.grad = None均无法彻底清除梯度,使用的Diffusers版本为0.3.0。
可行解决方案
1. 对整个Pipeline执行全局梯度清除
只清除Unet的梯度可能不够,因为Pipeline中包含多个子模块(如文本编码器、VAE等),部分梯度可能分散在这些模块中。直接对整个Pipeline调用zero_grad():
model.zero_grad(set_to_none=True)
set_to_none=True会直接将梯度置为None,比默认的set_to_none=False(将梯度置0)更节省显存,也更彻底。
2. 遍历所有子模块强制清除梯度
如果全局调用仍未解决问题,可以手动遍历Pipeline的所有子模块,逐个执行梯度清除:
for submodule in model.modules(): if hasattr(submodule, 'zero_grad'): submodule.zero_grad(set_to_none=True)
这种方式会覆盖Pipeline中所有带参数的子模块,确保没有遗漏任何梯度残留。
3. 注意事项
torch.cuda.empty_cache()仅用于释放未被PyTorch占用的显存,和梯度清除无关,无需在梯度清除代码中加入。- 若在循环中多次调用生成方法,务必在每次梯度编辑完成后立即执行上述清除操作,避免梯度跨步骤积累。
内容的提问来源于stack exchange,提问作者python_noob
相关产品推荐
相关产品推荐

