From 2edefa6859d4493ca7bd4b4064f12780a86f10dd Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=88=98=E9=9B=AA=E5=B3=B0?= Date: Sat, 18 Jan 2025 20:53:35 +0800 Subject: [PATCH] Optimize memory management --- PulidFluxHook.py | 18 +++++++++--------- README.md | 2 +- pulidflux.py | 29 +++++++++++++++++++++-------- pyproject.toml | 2 +- 4 files changed, 32 insertions(+), 19 deletions(-) diff --git a/PulidFluxHook.py b/PulidFluxHook.py index 606439f..2700791 100644 --- a/PulidFluxHook.py +++ b/PulidFluxHook.py @@ -6,16 +6,16 @@ import comfy from .patch_util import PatchKeys def set_model_dit_patch_replace(model, patch_kwargs, key): - to = model.model_options["transformer_options"].copy() + to = model.model_options["transformer_options"] if "patches_replace" not in to: to["patches_replace"] = {} else: - to["patches_replace"] = to["patches_replace"].copy() + to["patches_replace"] = to["patches_replace"] if "dit" not in to["patches_replace"]: to["patches_replace"]["dit"] = {} else: - to["patches_replace"]["dit"] = to["patches_replace"]["dit"].copy() + to["patches_replace"]["dit"] = to["patches_replace"]["dit"] if key not in to["patches_replace"]["dit"]: if "double_block" in key: @@ -25,12 +25,12 @@ def set_model_dit_patch_replace(model, patch_kwargs, key): to["patches_replace"]["dit"][key] = DitDoubleBlockReplace(pulid_patch, **patch_kwargs) else: to["patches_replace"]["dit"][key] = DitSingleBlockReplace(pulid_patch, **patch_kwargs) - model.model_options["transformer_options"] = to + # model.model_options["transformer_options"] = to else: to["patches_replace"]["dit"][key].add(pulid_patch, **patch_kwargs) def pulid_patch(img, pulid_model=None, ca_idx=None, weight=1.0, embedding=None, mask=None, transformer_options={}): - pulid_img = weight * pulid_model.pulid_ca[ca_idx].to(img.device)(embedding, img) + pulid_img = weight * pulid_model.model.pulid_ca[ca_idx](embedding, img) if mask is not None: pulid_temp_attrs = transformer_options.get(PatchKeys.pulid_patch_key_attrs, {}) latent_image_shape = pulid_temp_attrs.get("latent_image_shape", None) @@ -41,12 +41,12 @@ def pulid_patch(img, pulid_model=None, ca_idx=None, weight=1.0, embedding=None, mask = comfy.ldm.common_dit.pad_to_patch_size(mask, (patch_size, patch_size)) mask = rearrange(mask, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch_size, pw=patch_size) # (b, seq_len, _) =>(b, seq_len, seq_len) - mask = mask[..., 0].unsqueeze(-1).repeat(1, 1, mask.shape[1]) + mask = mask[..., 0].unsqueeze(-1).repeat(1, 1, mask.shape[1]).to(dtype=pulid_img.dtype) del patch_size, latent_image_shape pulid_img = pulid_img * mask - del mask + del mask, pulid_temp_attrs return pulid_img @@ -65,7 +65,7 @@ class DitDoubleBlockReplace: def __call__(self, input_args, extra_options): transformer_options = extra_options["transformer_options"] pulid_temp_attrs = transformer_options.get(PatchKeys.pulid_patch_key_attrs, {}) - sigma = pulid_temp_attrs["timesteps"].detach().cpu()[0] + sigma = pulid_temp_attrs["timesteps"].detach().cpu().item() out = extra_options["original_block"](input_args) img = out['img'] temp_img = img @@ -112,7 +112,7 @@ class DitSingleBlockReplace: out = extra_options["original_block"](input_args) - sigma = pulid_temp_attrs["timesteps"][0] + sigma = pulid_temp_attrs["timesteps"][0].detach().cpu().item() img = out['img'] txt = pulid_temp_attrs['double_blocks_txt'] real_img, txt = img[:, txt.shape[1]:, ...], img[:, :txt.shape[1], ...] diff --git a/README.md b/README.md index c1b1743..286b26e 100644 --- a/README.md +++ b/README.md @@ -39,7 +39,7 @@ Please see [ComfyUI-PuLID-Flux](https://github.com/balazik/ComfyUI-PuLID-Flux) - If you want use with [TeaCache](https://github.com/ali-vilab/TeaCache), must put it before node [`FluxForwardOverrider` and `ApplyTeaCachePatch`](https://github.com/lldacing/ComfyUI_Patches_ll). - If you want use with [Comfy-WaveSpeed](https://github.com/chengzeyi/Comfy-WaveSpeed), must put it before node `ApplyFBCacheOnModel`. - FixPulidFluxPatch (Deprecated) - - If you want use with [TeaCache](https://github.com/ali-vilab/TeaCache), must ~~link it after node `ApplyPulidFlux`, and~~ link node [`FluxForwardOverrider` and `ApplyTeaCachePatch`](https://github.com/lldacing/ComfyUI_Patches_ll) after it. + - If you want use with [TeaCache](https://github.com/ali-vilab/TeaCache), must link it after node `ApplyPulidFlux`, and link node [`FluxForwardOverrider` and `ApplyTeaCachePatch`](https://github.com/lldacing/ComfyUI_Patches_ll) after it. ## Thanks diff --git a/pulidflux.py b/pulidflux.py index 288b113..89f3810 100644 --- a/pulidflux.py +++ b/pulidflux.py @@ -12,6 +12,7 @@ from insightface.app import FaceAnalysis from facexlib.parsing import init_parsing_model from facexlib.utils.face_restoration_helper import FaceRestoreHelper +from comfy import model_management from .eva_clip.constants import OPENAI_DATASET_MEAN, OPENAI_DATASET_STD from .encoders_flux import IDFormer, PerceiverAttentionCA @@ -101,12 +102,18 @@ class PulidFluxModelLoader: model_path = folder_paths.get_full_path("pulid", pulid_file) # Also initialize the model, takes longer to load but then it doesn't have to be done every time you change parameters in the apply node + offload_device = model_management.unet_offload_device() + load_device = model_management.get_torch_device() + model = PulidFluxModel() logging.info("Loading PuLID-Flux model.") model.from_pretrained(path=model_path) - return (model,) + model_patcher = comfy.model_patcher.ModelPatcher(model, load_device=load_device, offload_device=offload_device) + del model + + return (model_patcher,) class PulidFluxInsightFaceLoader: @classmethod @@ -195,14 +202,15 @@ class ApplyPulidFlux: dtype = model.model.manual_cast_dtype eva_clip.to(device, dtype=dtype) - pulid_flux.to(device, dtype=dtype) + pulid_flux.model.to(dtype=dtype) + model_management.load_model_gpu(pulid_flux) if attn_mask is not None: if attn_mask.dim() > 3: attn_mask = attn_mask.squeeze(-1) elif attn_mask.dim() < 3: attn_mask = attn_mask.unsqueeze(0) - attn_mask = attn_mask.to(device, dtype=dtype) + # attn_mask = attn_mask.to(device, dtype=dtype) image = tensor_to_image(image) @@ -279,8 +287,9 @@ class ApplyPulidFlux: id_cond = torch.cat([iface_embeds, id_cond_vit], dim=-1) # Pulid_encoder - cond.append(pulid_flux.get_embeds(id_cond, id_vit_hidden)) + cond.append(pulid_flux.model.get_embeds(id_cond, id_vit_hidden)) + eva_clip.to(torch.device('cpu')) if not cond: # No faces detected, return the original model logging.warning("PuLID warning: No faces detected in any of the given images, returning unmodified model.") @@ -305,16 +314,19 @@ class ApplyPulidFlux: ca_idx = 0 for i in range(19): - if i % pulid_flux.double_interval == 0: + if i % pulid_flux.model.double_interval == 0: patch_kwargs["ca_idx"] = ca_idx set_model_dit_patch_replace(model, patch_kwargs, ("double_block", i)) ca_idx += 1 for i in range(38): - if i % pulid_flux.single_interval == 0: + if i % pulid_flux.model.single_interval == 0: patch_kwargs["ca_idx"] = ca_idx set_model_dit_patch_replace(model, patch_kwargs, ("single_block", i)) ca_idx += 1 + if len(model.get_additional_models_with_key("pulid_flux_model_patcher")) == 0: + model.set_additional_models("pulid_flux_model_patcher", [pulid_flux]) + if len(model.get_wrappers(comfy.patcher_extension.WrappersMP.OUTER_SAMPLE, wrappers_name)) == 0: # Just add it once when connecting in series model.add_wrapper_with_key(comfy.patcher_extension.WrappersMP.OUTER_SAMPLE, wrappers_name, pulid_outer_sample_wrappers_with_override) @@ -322,7 +334,7 @@ class ApplyPulidFlux: # Just add it once when connecting in series model.add_wrapper_with_key(comfy.patcher_extension.WrappersMP.APPLY_MODEL, wrappers_name, pulid_apply_model_wrappers) - del eva_clip, face_analysis, pulid_flux + del eva_clip, face_analysis, pulid_flux, face_helper, attn_mask return (model,) @@ -381,6 +393,7 @@ def pulid_outer_sample_wrappers_with_override(wrapper_executor, noise, latent_im finally: del PULID_model_patch['latent_image_shape'] clean_hook(diffusion_model) + del diffusion_model, cfg_guider return out @@ -405,7 +418,7 @@ def pulid_apply_model_wrappers(wrapper_executor, x, t, c_concat=None, c_crossatt finally: if PatchKeys.running_net_model in transformer_options: del transformer_options[PatchKeys.running_net_model] - del PULID_model_patch['timesteps'] + del PULID_model_patch['timesteps'], base_model return out diff --git a/pyproject.toml b/pyproject.toml index 4c4db89..f96b664 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "comfyui_pulid_flux_ll" description = "The implementation for PuLID-Flux, support use with TeaCache and WaveSpeed, no model pollution." -version = "1.0.4" +version = "1.0.5" license = {file = "LICENSE"} dependencies = ['facexlib', 'insightface', 'onnxruntime', 'onnxruntime-gpu; sys_platform != "darwin" and platform_machine == "x86_64"', 'ftfy', 'timm']