From 70c2e84d26f727dabc84e2b7eee07c4bdc80bcec Mon Sep 17 00:00:00 2001 From: Disty0 Date: Mon, 12 Aug 2024 01:02:13 +0300 Subject: [PATCH] Prompt cache support for Flux --- modules/model_t5.py | 4 ++-- modules/processing_args.py | 8 +++++++- modules/prompt_parser_diffusers.py | 15 ++++++++++++--- 3 files changed, 21 insertions(+), 6 deletions(-) diff --git a/modules/model_t5.py b/modules/model_t5.py index c425d9d7a..e5201e675 100644 --- a/modules/model_t5.py +++ b/modules/model_t5.py @@ -86,7 +86,7 @@ def set_t5(pipe, module, t5=None, cache_dir=None): elif shared.opts.diffusers_offload_mode == "cpu": if not hasattr(pipe, "_all_hooks") or len(pipe._all_hooks) == 0: # pylint: disable=protected-access pipe.enable_model_cpu_offload(device=devices.device) - else: - pipe.maybe_free_model_hooks() + if hasattr(pipe, "maybe_free_model_hooks"): + pipe.maybe_free_model_hooks() devices.torch_gc() return pipe diff --git a/modules/processing_args.py b/modules/processing_args.py index 91f10c80d..d08783256 100644 --- a/modules/processing_args.py +++ b/modules/processing_args.py @@ -110,7 +110,11 @@ def set_pipeline_args(p, model, prompts: list, negative_prompts: list, prompts_2 shared.log.error(f'Sampler timesteps: {e}') else: shared.log.warning(f'Sampler: sampler={model.scheduler.__class__.__name__} timesteps not supported') - if shared.opts.prompt_attention != 'Fixed attention' and ('StableDiffusion' in model.__class__.__name__ or 'StableCascade' in model.__class__.__name__) and 'Onnx' not in model.__class__.__name__: + if shared.opts.prompt_attention != 'Fixed attention' and 'Onnx' not in model.__class__.__name__ and ( + 'StableDiffusion' in model.__class__.__name__ or + 'StableCascade' in model.__class__.__name__ or + 'Flux' in model.__class__.__name__ + ): try: prompt_parser_diffusers.encode_prompts(model, p, prompts, negative_prompts, steps=steps, clip_skip=clip_skip) parser = shared.opts.prompt_attention @@ -132,6 +136,8 @@ def set_pipeline_args(p, model, prompts: list, negative_prompts: list, prompts_2 args['pooled_prompt_embeds'] = p.positive_pooleds[0] elif 'StableDiffusion3' in model.__class__.__name__ and len(getattr(p, 'positive_pooleds', [])) > 0: args['pooled_prompt_embeds'] = p.positive_pooleds[0] + elif 'Flux' in model.__class__.__name__ and len(getattr(p, 'positive_pooleds', [])) > 0: + args['pooled_prompt_embeds'] = p.positive_pooleds[0] else: args['prompt'] = prompts if 'negative_prompt' in possible: diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index a4ba0597c..ae4405a45 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -147,7 +147,12 @@ def get_tokens(msg, prompt): def encode_prompts(pipe, p, prompts: list, negative_prompts: list, steps: int, clip_skip: typing.Optional[int] = None): - if 'StableDiffusion' not in pipe.__class__.__name__ and 'DemoFusion' not in pipe.__class__.__name__ and 'StableCascade' not in pipe.__class__.__name__: + if ( + 'StableDiffusion' not in pipe.__class__.__name__ and + 'DemoFusion' not in pipe.__class__.__name__ and + 'StableCascade' not in pipe.__class__.__name__ and + 'Flux' not in pipe.__class__.__name__ + ): shared.log.warning(f"Prompt parser not supported: {pipe.__class__.__name__}") return elif shared.opts.sd_textencoder_cache and prompts == cache.get('prompts', None) and negative_prompts == cache.get('negative_prompts', None) and clip_skip == cache.get('clip_skip', None) and cache.get('model_type', None) == shared.sd_model_type and steps == cache.get('steps', None): @@ -168,7 +173,7 @@ def encode_prompts(pipe, p, prompts: list, negative_prompts: list, steps: int, c p.negative_embeds = [] p.negative_pooleds = [] - if shared.opts.diffusers_offload_mode in {"balanced", "cpu"} and hasattr(pipe, "_all_hooks") and hasattr(pipe, "maybe_free_model_hooks"): + if hasattr(pipe, "maybe_free_model_hooks"): # if the last job is interrupted, model will stay in the vram and cause oom, send everything back to cpu before continuing pipe.maybe_free_model_hooks() devices.torch_gc() @@ -204,7 +209,7 @@ def encode_prompts(pipe, p, prompts: list, negative_prompts: list, steps: int, c if debug_enabled: get_tokens('positive', prompts[0]) get_tokens('negative', negative_prompts[0]) - if shared.opts.diffusers_offload_mode in {"balanced", "cpu"} and hasattr(pipe, "_all_hooks") and hasattr(pipe, "maybe_free_model_hooks"): + if hasattr(pipe, "maybe_free_model_hooks"): # text encoder will stay in the vram and cause oom, send everything back to cpu before continuing pipe.maybe_free_model_hooks() debug(f"Prompt encode: time={(time.time() - t0):.3f}") @@ -332,6 +337,10 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c negatives.pop(0) negative_weights.pop(0) + if "Flux" in pipe.__class__.__name__: # clip is only used for the pooled embeds + prompt_embeds, pooled_prompt_embeds, _ = pipe.encode_prompt(prompt=prompt, prompt_2=prompt_2, device=device, num_images_per_prompt=1) + return prompt_embeds, pooled_prompt_embeds, None, None # no negative support + embedding_providers = prepare_embedding_providers(pipe, clip_skip) empty_embedding_providers = None if 'StableCascade' in pipe.__class__.__name__: