support for sdxl-inpaint model

This commit is contained in:
wangqyqq
2023-12-21 20:15:51 +08:00
parent cf2772fab0
commit 9feb034e34
4 changed files with 127 additions and 1 deletions

View File

@@ -34,6 +34,11 @@ def get_learned_conditioning(self: sgm.models.diffusion.DiffusionEngine, batch:
def apply_model(self: sgm.models.diffusion.DiffusionEngine, x, t, cond):
sd = self.model.state_dict()
diffusion_model_input = sd.get('diffusion_model.input_blocks.0.0.weight', None)
if diffusion_model_input.shape[1] == 9:
x = torch.cat([x] + cond['c_concat'], dim=1)
return self.model(x, t, cond)