De-dupe Cascade

This commit is contained in:
Disty0
2024-08-01 18:12:02 +03:00
parent b50a8601fe
commit 9965ef75e7
3 changed files with 28 additions and 10 deletions
+8
View File
@@ -120,6 +120,14 @@ def load_cascade_combined(checkpoint_info, diffusers_load_config):
sd_model.decoder_pipe.text_encoder = sd_model.text_encoder = None # Nothing uses the decoder's text encoder
sd_model.prior_pipe.image_encoder = sd_model.prior_image_encoder = None # No img2img is implemented yet
sd_model.prior_pipe.feature_extractor = sd_model.prior_feature_extractor = None # No img2img is implemented yet
#de-dupe
del sd_model.decoder_pipe.text_encoder
del sd_model.prior_prior
del sd_model.prior_text_encoder
del sd_model.prior_tokenizer
del sd_model.prior_scheduler
del sd_model.prior_feature_extractor
del sd_model.prior_image_encoder
shared.log.debug(f'StableCascade combined: {sd_model.__class__.__name__}')
return sd_model
+16 -10
View File
@@ -36,13 +36,12 @@ def apply_compile_to_model(sd_model, function, options, op=None):
if hasattr(sd_model, 'decoder_pipe') and hasattr(sd_model, 'decoder'):
sd_model.decoder = None
sd_model.decoder = sd_model.decoder_pipe.decoder = function(sd_model.decoder_pipe.decoder)
if hasattr(sd_model, 'prior_pipe') and hasattr(sd_model, 'prior_prior'):
sd_model.prior_prior = None
if hasattr(sd_model, 'prior_pipe') and hasattr(sd_model.prior_pipe, 'prior'):
if op == "nncf" and "StableCascade" in sd_model.__class__.__name__: # fixes dtype errors
backup_clip_txt_pooled_mapper = copy.deepcopy(sd_model.prior_pipe.prior.clip_txt_pooled_mapper)
sd_model.prior_prior = sd_model.prior_pipe.prior = function(sd_model.prior_pipe.prior)
sd_model.prior_pipe.prior = function(sd_model.prior_pipe.prior)
if op == "nncf" and "StableCascade" in sd_model.__class__.__name__:
sd_model.prior_prior.clip_txt_pooled_mapper = sd_model.prior_pipe.prior.clip_txt_pooled_mapper = backup_clip_txt_pooled_mapper
sd_model.prior_pipe.prior.clip_txt_pooled_mapper = backup_clip_txt_pooled_mapper
if "VAE" in options:
if hasattr(sd_model, 'vae') and hasattr(sd_model.vae, 'decode'):
sd_model.vae = function(sd_model.vae)
@@ -50,13 +49,14 @@ def apply_compile_to_model(sd_model, function, options, op=None):
sd_model.movq = function(sd_model.movq)
if hasattr(sd_model, 'vqgan') and hasattr(sd_model.vqgan, 'decode'):
sd_model.vqgan = function(sd_model.vqgan)
if hasattr(sd_model, 'decoder_pipe') and hasattr(sd_model.decoder_pipe, 'vqgan'):
sd_model.decoder_pipe.vqgan = sd_model.vqgan
if hasattr(sd_model, 'image_encoder') and hasattr(sd_model.image_encoder, 'config'):
sd_model.image_encoder = function(sd_model.image_encoder)
if "Text Encoder" in options:
if hasattr(sd_model, 'text_encoder') and hasattr(sd_model.text_encoder, 'config'):
if hasattr(sd_model, 'decoder_pipe'):
sd_model.text_encoder = None
sd_model.text_encoder = sd_model.decoder_pipe.text_encoder = function(sd_model.decoder_pipe.text_encoder)
if hasattr(sd_model, 'decoder_pipe') and hasattr(sd_model.decoder_pipe, 'text_encoder'):
sd_model.decoder_pipe.text_encoder = function(sd_model.decoder_pipe.text_encoder)
else:
if op == "nncf" and sd_model.text_encoder.__class__.__name__ in {"T5EncoderModel", "UMT5EncoderModel"}:
from modules.sd_hijack import NNCF_T5DenseGatedActDense # T5DenseGatedActDense uses fp32
@@ -67,6 +67,13 @@ def apply_compile_to_model(sd_model, function, options, op=None):
)
sd_model.text_encoder = function(sd_model.text_encoder)
if hasattr(sd_model, 'text_encoder_2') and hasattr(sd_model.text_encoder_2, 'config'):
if op == "nncf" and sd_model.text_encoder_2.__class__.__name__ in {"T5EncoderModel", "UMT5EncoderModel"}:
from modules.sd_hijack import NNCF_T5DenseGatedActDense # T5DenseGatedActDense uses fp32
for i in range(len(sd_model.text_encoder_2.encoder.block)):
sd_model.text_encoder_2.encoder.block[i].layer[1].DenseReluDense = NNCF_T5DenseGatedActDense(
sd_model.text_encoder_2.encoder.block[i].layer[1].DenseReluDense,
dtype=torch.float32 if devices.dtype != torch.bfloat16 else torch.bfloat16
)
sd_model.text_encoder_2 = function(sd_model.text_encoder_2)
if hasattr(sd_model, 'text_encoder_3') and hasattr(sd_model.text_encoder_3, 'config'):
if op == "nncf" and sd_model.text_encoder_3.__class__.__name__ in {"T5EncoderModel", "UMT5EncoderModel"}:
@@ -77,9 +84,8 @@ def apply_compile_to_model(sd_model, function, options, op=None):
dtype=torch.float32 if devices.dtype != torch.bfloat16 else torch.bfloat16
)
sd_model.text_encoder_3 = function(sd_model.text_encoder_3)
if hasattr(sd_model, 'prior_pipe') and hasattr(sd_model, 'prior_text_encoder'):
sd_model.prior_text_encoder = None
sd_model.prior_text_encoder = sd_model.prior_pipe.text_encoder = function(sd_model.prior_pipe.text_encoder)
if hasattr(sd_model, 'prior_pipe') and hasattr(sd_model.prior_pipe, 'text_encoder'):
sd_model.prior_pipe.text_encoder = function(sd_model.prior_pipe.text_encoder)
return sd_model
+4
View File
@@ -50,6 +50,8 @@ def create_sampler(name, model):
if name == 'Default' and hasattr(model, 'scheduler'):
if getattr(model, "default_scheduler", None) is not None:
model.scheduler = copy.deepcopy(model.default_scheduler)
if hasattr(model, "prior_pipe") and hasattr(model.prior_pipe, "scheduler"):
model.prior_pipe.scheduler = copy.deepcopy(model.default_scheduler)
config = {k: v for k, v in model.scheduler.config.items() if not k.startswith('_')}
shared.log.debug(f'Sampler default {type(model.scheduler).__name__}: {config}')
return model.scheduler
@@ -69,6 +71,8 @@ def create_sampler(name, model):
if not hasattr(model, 'scheduler_config'):
model.scheduler_config = sampler.sampler.config.copy()
model.scheduler = sampler.sampler
if hasattr(model, "prior_pipe") and hasattr(model.prior_pipe, "scheduler"):
model.prior_pipe.scheduler = sampler.sampler
shared.log.debug(f'Sampler: sampler="{sampler.name}" config={sampler.config}')
return sampler.sampler
else: