Optimize memory management

This commit is contained in:
刘雪峰
2025-01-18 20:53:35 +08:00
parent 60797583e6
commit 2edefa6859
4 changed files with 32 additions and 19 deletions
+9 -9
View File
@@ -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], ...]
+1 -1
View File
@@ -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
+21 -8
View File
@@ -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
+1 -1
View File
@@ -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']