sd_hijack_checkpoint.py 1.2 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546
  1. from torch.utils.checkpoint import checkpoint
  2. import ldm.modules.attention
  3. import ldm.modules.diffusionmodules.openaimodel
  4. def BasicTransformerBlock_forward(self, x, context=None):
  5. return checkpoint(self._forward, x, context)
  6. def AttentionBlock_forward(self, x):
  7. return checkpoint(self._forward, x)
  8. def ResBlock_forward(self, x, emb):
  9. return checkpoint(self._forward, x, emb)
  10. stored = []
  11. def add():
  12. if len(stored) != 0:
  13. return
  14. stored.extend([
  15. ldm.modules.attention.BasicTransformerBlock.forward,
  16. ldm.modules.diffusionmodules.openaimodel.ResBlock.forward,
  17. ldm.modules.diffusionmodules.openaimodel.AttentionBlock.forward
  18. ])
  19. ldm.modules.attention.BasicTransformerBlock.forward = BasicTransformerBlock_forward
  20. ldm.modules.diffusionmodules.openaimodel.ResBlock.forward = ResBlock_forward
  21. ldm.modules.diffusionmodules.openaimodel.AttentionBlock.forward = AttentionBlock_forward
  22. def remove():
  23. if len(stored) == 0:
  24. return
  25. ldm.modules.attention.BasicTransformerBlock.forward = stored[0]
  26. ldm.modules.diffusionmodules.openaimodel.ResBlock.forward = stored[1]
  27. ldm.modules.diffusionmodules.openaimodel.AttentionBlock.forward = stored[2]
  28. stored.clear()