From 096fd04f53d1c9c9c174e934e190ff64082d0757 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Thu, 23 Nov 2023 09:42:24 -0500 Subject: [PATCH] fix model offload on long prompts and update requirements --- CHANGELOG.md | 2 +- modules/interrogate.py | 36 +----------------------------- modules/prompt_parser_diffusers.py | 8 +++---- modules/sd_models.py | 1 - requirements.txt | 4 ++-- 5 files changed, 7 insertions(+), 44 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 2f1e330cc..0c6d98907 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,6 +1,6 @@ # Change Log for SD.Next -## Update for 2023-11-22 +## Update for 2023-11-23 - **Diffusers** - **LCM** support for any *SD 1.5* or *SD-XL* model! diff --git a/modules/interrogate.py b/modules/interrogate.py index 76685dae9..5e7d249de 100644 --- a/modules/interrogate.py +++ b/modules/interrogate.py @@ -12,27 +12,23 @@ from modules import devices, paths, shared, lowvram, modelloader, errors blip_image_eval_size = 384 clip_model_name = 'ViT-L/14' - Category = namedtuple("Category", ["name", "topn", "items"]) - re_topn = re.compile(r"\.top(\d+)\.") + def category_types(): return [f.stem for f in Path(shared.interrogator.content_dir).glob('*.txt')] def download_default_clip_interrogate_categories(content_dir): shared.log.info("Downloading CLIP categories...") - tmpdir = f"{content_dir}_tmp" cat_types = ["artists", "flavors", "mediums", "movements"] - try: os.makedirs(tmpdir, exist_ok=True) for category_type in cat_types: torch.hub.download_url_to_file(f"https://raw.githubusercontent.com/pharmapsychotic/clip-interrogator/main/clip_interrogator/data/{category_type}.txt", os.path.join(tmpdir, f"{category_type}.txt")) os.rename(tmpdir, content_dir) - except Exception as e: errors.display(e, "downloading default CLIP interrogate categories") finally: @@ -56,10 +52,8 @@ class InterrogateModels: def categories(self): if not os.path.exists(self.content_dir): download_default_clip_interrogate_categories(self.content_dir) - if self.loaded_categories is not None and self.skip_categories == shared.opts.interrogate_clip_skip_categories: return self.loaded_categories - self.loaded_categories = [] if os.path.exists(self.content_dir): @@ -74,14 +68,12 @@ class InterrogateModels: with open(filename, "r", encoding="utf8") as file: lines = [x.strip() for x in file.readlines()] self.loaded_categories.append(Category(name=filename.stem, topn=topn, items=lines)) - return self.loaded_categories def create_fake_fairscale(self): class FakeFairscale: def checkpoint_wrapper(self): pass - sys.modules["fairscale.nn.checkpoint.checkpoint_activations"] = FakeFairscale def load_blip_model(self): @@ -95,7 +87,6 @@ class InterrogateModels: ext_filter=[".pth"], download_name='model_base_caption_capfilt_large.pth', ) - blip_model = models.blip.blip_decoder(pretrained=files[0], image_size=blip_image_eval_size, vit='base', med_config=os.path.join(paths.paths["BLIP"], "configs", "med_config.json")) # pylint: disable=c-extension-no-member blip_model.eval() @@ -103,15 +94,12 @@ class InterrogateModels: def load_clip_model(self): import clip - if self.running_on_cpu: model, preprocess = clip.load(clip_model_name, device="cpu", download_root=shared.opts.clip_models_path) else: model, preprocess = clip.load(clip_model_name, download_root=shared.opts.clip_models_path) - model.eval() model = model.to(devices.device_interrogate) - return model, preprocess def load(self): @@ -119,16 +107,12 @@ class InterrogateModels: self.blip_model = self.load_blip_model() if not shared.opts.no_half and not self.running_on_cpu: self.blip_model = self.blip_model.half() - self.blip_model = self.blip_model.to(devices.device_interrogate) - if self.clip_model is None: self.clip_model, self.clip_preprocess = self.load_clip_model() if not shared.opts.no_half and not self.running_on_cpu: self.clip_model = self.clip_model.half() - self.clip_model = self.clip_model.to(devices.device_interrogate) - self.dtype = next(self.clip_model.parameters()).dtype def send_clip_to_ram(self): @@ -144,27 +128,21 @@ class InterrogateModels: def unload(self): self.send_clip_to_ram() self.send_blip_to_ram() - devices.torch_gc() def rank(self, image_features, text_array, top_count=1): import clip - devices.torch_gc() - if shared.opts.interrogate_clip_dict_limit != 0: text_array = text_array[0:int(shared.opts.interrogate_clip_dict_limit)] - top_count = min(top_count, len(text_array)) text_tokens = clip.tokenize(list(text_array), truncate=True).to(devices.device_interrogate) text_features = self.clip_model.encode_text(text_tokens).type(self.dtype) text_features /= text_features.norm(dim=-1, keepdim=True) - similarity = torch.zeros((1, len(text_array))).to(devices.device_interrogate) for i in range(image_features.shape[0]): similarity += (100.0 * image_features[i].unsqueeze(0) @ text_features.T).softmax(dim=-1) similarity /= image_features.shape[0] - top_probs, top_labels = similarity.cpu().topk(top_count, dim=-1) return [(text_array[top_labels[0][i].numpy()], (top_probs[0][i].numpy()*100)) for i in range(top_count)] @@ -174,10 +152,8 @@ class InterrogateModels: transforms.ToTensor(), transforms.Normalize((0.48145466, 0.4578275, 0.40821073), (0.26862954, 0.26130258, 0.27577711)) ])(pil_image).unsqueeze(0).type(self.dtype).to(devices.device_interrogate) - with devices.inference_context(): caption = self.blip_model.generate(gpu_image, sample=False, num_beams=shared.opts.interrogate_clip_num_beams, min_length=shared.opts.interrogate_clip_min_length, max_length=shared.opts.interrogate_clip_max_length) - return caption[0] def interrogate(self, pil_image): @@ -187,22 +163,15 @@ class InterrogateModels: if shared.cmd_opts.lowvram or shared.cmd_opts.medvram: lowvram.send_everything_to_cpu() devices.torch_gc() - self.load() - caption = self.generate_caption(pil_image) self.send_blip_to_ram() devices.torch_gc() - res = caption - clip_image = self.clip_preprocess(pil_image).unsqueeze(0).type(self.dtype).to(devices.device_interrogate) - with devices.inference_context(), devices.autocast(): image_features = self.clip_model.encode_image(clip_image).type(self.dtype) - image_features /= image_features.norm(dim=-1, keepdim=True) - for _name, topn, items in self.categories(): matches = self.rank(image_features, items, top_count=topn) for match, score in matches: @@ -210,12 +179,9 @@ class InterrogateModels: res += f", ({match}:{score/100:.3f})" else: res += f", {match}" - except Exception as e: errors.display(e, 'interrogate') res += "" - self.unload() shared.state.end() - return res diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index b9fbf8908..124bbb41d 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -69,7 +69,7 @@ def encode_prompts(pipeline, prompts: list, negative_prompts: list, clip_skip: t negative_embeds = [] negative_pooleds = [] for i in range(len(prompts)): - prompt_embed, positive_pooled, negative_embed, negative_pooled = get_weighted_text_embeddings(pipeline,prompts[i], negative_prompts[i], clip_skip) + prompt_embed, positive_pooled, negative_embed, negative_pooled = get_weighted_text_embeddings(pipeline, prompts[i], negative_prompts[i], clip_skip) prompt_embeds.append(prompt_embed) positive_pooleds.append(positive_pooled) negative_embeds.append(negative_embed) @@ -119,8 +119,7 @@ def pad_to_same_length(embeds): 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_batched = torch.cat([empty_embed[0]] * embeds[0].shape[0]) + 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): while embed.shape[1] < max_token_count: @@ -152,7 +151,6 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c negative_prompt_embeds = [] pooled_prompt_embeds = None negative_pooled_prompt_embeds = None - for i in range(len(embedding_providers)): # add BREAK keyword that splits the prompt into multiple fragments text = positives[i] @@ -168,7 +166,7 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c 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=pipe.device, should_return_tokens=True) negative_prompt_embeds.append(embed) if prompt_embeds[-1].shape[-1] > 768: diff --git a/modules/sd_models.py b/modules/sd_models.py index dc74a17c1..dccba9fe8 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -1106,7 +1106,6 @@ def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None, shared.log.debug(f"Model created from config: {checkpoint_config}") sd_model.used_config = checkpoint_config sd_model.has_accelerate = False - sd_model.is_sdxl = False # a1111 compatibility item timer.record("create") ok = load_model_weights(sd_model, checkpoint_info, state_dict, timer) if not ok: diff --git a/requirements.txt b/requirements.txt index 7e5eaefc0..bbe4874e4 100644 --- a/requirements.txt +++ b/requirements.txt @@ -49,7 +49,7 @@ clip-interrogator==0.6.0 antlr4-python3-runtime==4.9.3 requests==2.31.0 tqdm==4.66.1 -accelerate==0.20.3 +accelerate==0.24.1 opencv-python-headless==4.7.0.72 diffusers==0.23.1 einops==0.4.1 @@ -61,7 +61,7 @@ numba==0.57.1 pandas==1.5.3 protobuf==3.20.3 pytorch_lightning==1.9.4 -transformers==4.35.1 +transformers==4.35.2 tomesd==0.1.3 urllib3==1.26.15 Pillow==10.1.0