diff options
author | wangqyqq <wangqyqq@163.com> | 2023-12-21 20:15:51 +0800 |
---|---|---|
committer | wangqyqq <wangqyqq@163.com> | 2023-12-21 20:15:51 +0800 |
commit | 9feb034e343d6d7ef63395821658fb3774b30a24 (patch) | |
tree | dc4ef890656b264d3db1ff671fa3126458b09f10 /modules/sd_models_xl.py | |
parent | cf2772fab0af5573da775e7437e6acdca424f26e (diff) |
support for sdxl-inpaint model
Diffstat (limited to 'modules/sd_models_xl.py')
-rw-r--r-- | modules/sd_models_xl.py | 5 |
1 files changed, 5 insertions, 0 deletions
diff --git a/modules/sd_models_xl.py b/modules/sd_models_xl.py index 01123321..d8a9a73b 100644 --- a/modules/sd_models_xl.py +++ b/modules/sd_models_xl.py @@ -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)
|