Optimize memory management
This commit is contained in:
+9
-9
@@ -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], ...]
|
||||
|
||||
@@ -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
@@ -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
@@ -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']
|
||||
|
||||
|
||||
Reference in New Issue
Block a user