From 567b9dd4359610445ba4186d3e413b36f7f6764e Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Thu, 23 Nov 2023 10:31:46 -0500 Subject: [PATCH] offloading deal with meta tensors --- CHANGELOG.md | 2 +- installer.py | 2 +- modules/prompt_parser_diffusers.py | 31 ++++++++++++++++-------------- 3 files changed, 19 insertions(+), 16 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 0c6d98907..a836912cd 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -19,6 +19,7 @@ *Hint*: **70+ it/s** is possible on *RTX4090* with no special tweaks - Add additional pipeline types for manual model loads when loading from `safetensors` - Updated logic for calculating **steps** when using base/hires/refiner workflows + - Improve **model offloading** for both model and sequential cpu offload when dealing with meta tensors - Safe model offloading for non-standard models - Fix **DPM SDE** scheduler - Better support for SD 1.5 **inpainting** models @@ -26,7 +27,6 @@ - Enhance prompt parsing with long prompts and support for *BREAK* keyword Change-in-behavior: new line in prompt now means *BREAK* - Add alternative Lora loading algorithm, triggered if `SD_LORA_DIFFUSERS` is set - - Update to `diffusers==0.23.0` - **Models** - **Model merge** - completely redesigned, now based on best-of-class `meh` by @s1dlx diff --git a/installer.py b/installer.py index d81f9a141..e5aaeceed 100644 --- a/installer.py +++ b/installer.py @@ -190,7 +190,7 @@ def installed(package, friendly: str = None, reload = False, quiet = False): if args.experimental: log.warning(f"Package allowing experimental: {p[0]} {package_version} required {p[1]}") else: - log.warning(f"Package wrong version: {p[0]} {package_version} required {p[1]}") + log.warning(f"Package version mismatch: {p[0]} {package_version} required {p[1]}") else: if not quiet: log.debug(f"Package not found: {p[0]}") diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index 124bbb41d..735cf995a 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -3,7 +3,7 @@ import typing import torch from compel import ReturnedEmbeddingsType from compel.embeddings_provider import BaseTextualInversionManager, EmbeddingsProvider -from modules import shared, prompt_parser +from modules import shared, prompt_parser, devices debug = shared.log.info if os.environ.get('SD_PROMPT_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -97,6 +97,7 @@ def get_prompts_with_weights(prompt: str): def prepare_embedding_providers(pipe, clip_skip): + device = pipe.device if str(pipe.device) != 'meta' else devices.device embeddings_providers = [] if 'XL' in pipe.__class__.__name__: embedding_type = ReturnedEmbeddingsType.PENULTIMATE_HIDDEN_STATES_NON_NORMALIZED @@ -106,19 +107,20 @@ def prepare_embedding_providers(pipe, clip_skip): clip_skip = 2 embedding_type = CLIP_SKIP_MAPPING[clip_skip] if getattr(pipe, "tokenizer", None) is not None and getattr(pipe, "text_encoder", None) is not None: - embedding = EmbeddingsProvider(tokenizer=pipe.tokenizer, text_encoder=pipe.text_encoder, truncate=False, returned_embeddings_type=embedding_type, device=pipe.device) + embedding = EmbeddingsProvider(tokenizer=pipe.tokenizer, text_encoder=pipe.text_encoder, truncate=False, returned_embeddings_type=embedding_type, device=device) embeddings_providers.append(embedding) if getattr(pipe, "tokenizer_2", None) is not None and getattr(pipe, "text_encoder_2", None) is not None: - embedding = EmbeddingsProvider(tokenizer=pipe.tokenizer_2, text_encoder=pipe.text_encoder_2, truncate=False, returned_embeddings_type=embedding_type, device=pipe.device) + embedding = EmbeddingsProvider(tokenizer=pipe.tokenizer_2, text_encoder=pipe.text_encoder_2, truncate=False, returned_embeddings_type=embedding_type, device=device) embeddings_providers.append(embedding) return embeddings_providers -def pad_to_same_length(embeds): +def pad_to_same_length(pipe, embeds): + device = pipe.device if str(pipe.device) != 'meta' else devices.device try: #SDXL empty_embed = shared.sd_model.encode_prompt("") except Exception: #SD1.5 - empty_embed = shared.sd_model.encode_prompt("",shared.sd_model.device, 1, False) + empty_embed = shared.sd_model.encode_prompt("", device, 1, False) empty_batched = torch.cat([empty_embed[0].to(embeds[0].device)] * embeds[0].shape[0]) max_token_count = max([embed.shape[1] for embed in embeds]) for i, embed in enumerate(embeds): @@ -129,6 +131,7 @@ def pad_to_same_length(embeds): def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", clip_skip: int = None): + device = pipe.device if str(pipe.device) != 'meta' else devices.device prompt_2 = prompt.split("TE2:")[-1] neg_prompt_2 = neg_prompt.split("TE2:")[-1] prompt = prompt.split("TE2:")[0] @@ -160,35 +163,35 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c provider_embed = [] while 'BREAK' in text: pos = text.index('BREAK') - embed, ptokens = embedding_providers[i].get_embeddings_for_weighted_prompt_fragments(text_batch=[text[:pos]], fragment_weights_batch=[weights[:pos]], device=pipe.device, should_return_tokens=True) + embed, ptokens = embedding_providers[i].get_embeddings_for_weighted_prompt_fragments(text_batch=[text[:pos]], fragment_weights_batch=[weights[:pos]], device=device, should_return_tokens=True) provider_embed.append(embed) text = text[pos+1:] weights = weights[pos+1:] prompt_embeds.append(torch.cat(provider_embed, dim=1)) # negative prompt has no keywords - embed, ntokens = embedding_providers[i].get_embeddings_for_weighted_prompt_fragments(text_batch=[negatives[i]], fragment_weights_batch=[negative_weights[i]], device=pipe.device, should_return_tokens=True) + embed, ntokens = embedding_providers[i].get_embeddings_for_weighted_prompt_fragments(text_batch=[negatives[i]], fragment_weights_batch=[negative_weights[i]], device=device, should_return_tokens=True) negative_prompt_embeds.append(embed) if prompt_embeds[-1].shape[-1] > 768: if shared.opts.diffusers_pooled == "weighted": pooled_prompt_embeds = prompt_embeds[-1][ - torch.arange(prompt_embeds[-1].shape[0], device=pipe.device), - (ptokens.to(dtype=torch.int, device=pipe.device) == 49407) + torch.arange(prompt_embeds[-1].shape[0], device=device), + (ptokens.to(dtype=torch.int, device=device) == 49407) .int() .argmax(dim=-1), ] negative_pooled_prompt_embeds = negative_prompt_embeds[-1][ - torch.arange(negative_prompt_embeds[-1].shape[0], device=pipe.device), - (ntokens.to(dtype=torch.int, device=pipe.device) == 49407) + torch.arange(negative_prompt_embeds[-1].shape[0], device=device), + (ntokens.to(dtype=torch.int, device=device) == 49407) .int() .argmax(dim=-1), ] else: - pooled_prompt_embeds = embedding_providers[-1].get_pooled_embeddings(texts=[prompt_2], device=pipe.device) if prompt_embeds[-1].shape[-1] > 768 else None - negative_pooled_prompt_embeds = embedding_providers[-1].get_pooled_embeddings(texts=[neg_prompt_2], device=pipe.device) if negative_prompt_embeds[-1].shape[-1] > 768 else None + pooled_prompt_embeds = embedding_providers[-1].get_pooled_embeddings(texts=[prompt_2], device=device) if prompt_embeds[-1].shape[-1] > 768 else None + negative_pooled_prompt_embeds = embedding_providers[-1].get_pooled_embeddings(texts=[neg_prompt_2], device=device) if negative_prompt_embeds[-1].shape[-1] > 768 else None prompt_embeds = torch.cat(prompt_embeds, dim=-1) if len(prompt_embeds) > 1 else prompt_embeds[0] negative_prompt_embeds = torch.cat(negative_prompt_embeds, dim=-1) if len(negative_prompt_embeds) > 1 else negative_prompt_embeds[0] if prompt_embeds.shape[1] != negative_prompt_embeds.shape[1]: - [prompt_embeds, negative_prompt_embeds] = pad_to_same_length([prompt_embeds, negative_prompt_embeds]) + [prompt_embeds, negative_prompt_embeds] = pad_to_same_length(pipe, [prompt_embeds, negative_prompt_embeds]) return prompt_embeds, pooled_prompt_embeds, negative_prompt_embeds, negative_pooled_prompt_embeds