refactor ipadapter loading
This commit is contained in:
@@ -48,6 +48,9 @@ script_directory = os.path.dirname(os.path.abspath(__file__))
|
|||||||
IS_MODEL_CPU_OFFLOAD_ENABLED = False
|
IS_MODEL_CPU_OFFLOAD_ENABLED = False
|
||||||
|
|
||||||
class brushnet_model_loader:
|
class brushnet_model_loader:
|
||||||
|
# @classmethod
|
||||||
|
# def IS_CHANGED(s):
|
||||||
|
# return ""
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
return {"required": {
|
return {"required": {
|
||||||
@@ -62,6 +65,9 @@ class brushnet_model_loader:
|
|||||||
"default": "brushnet_segmentation_mask"
|
"default": "brushnet_segmentation_mask"
|
||||||
}),
|
}),
|
||||||
},
|
},
|
||||||
|
"optional": {
|
||||||
|
"ip_adapter": ("DIFFUSERSIPADAPTER",),
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("BRUSHNET",)
|
RETURN_TYPES = ("BRUSHNET",)
|
||||||
@@ -69,7 +75,7 @@ class brushnet_model_loader:
|
|||||||
FUNCTION = "loadmodel"
|
FUNCTION = "loadmodel"
|
||||||
CATEGORY = "BrushNetWrapper"
|
CATEGORY = "BrushNetWrapper"
|
||||||
|
|
||||||
def loadmodel(self, model, clip, vae, brushnet_model):
|
def loadmodel(self, model, clip, vae, brushnet_model, ip_adapter=None):
|
||||||
mm.soft_empty_cache()
|
mm.soft_empty_cache()
|
||||||
dtype = mm.unet_dtype()
|
dtype = mm.unet_dtype()
|
||||||
device = mm.get_torch_device()
|
device = mm.get_torch_device()
|
||||||
@@ -78,9 +84,10 @@ class brushnet_model_loader:
|
|||||||
"model": model,
|
"model": model,
|
||||||
"vae": vae,
|
"vae": vae,
|
||||||
"clip": clip,
|
"clip": clip,
|
||||||
"brushnet_model": brushnet_model
|
"brushnet_model": brushnet_model,
|
||||||
|
"ip_adapter": ip_adapter
|
||||||
}
|
}
|
||||||
if not hasattr(self, "model") or self.model == None or custom_config != self.current_config:
|
if not hasattr(self, "pipe") or custom_config != self.current_config:
|
||||||
global IS_MODEL_CPU_OFFLOAD_ENABLED
|
global IS_MODEL_CPU_OFFLOAD_ENABLED
|
||||||
IS_MODEL_CPU_OFFLOAD_ENABLED = False
|
IS_MODEL_CPU_OFFLOAD_ENABLED = False
|
||||||
pbar = comfy.utils.ProgressBar(5)
|
pbar = comfy.utils.ProgressBar(5)
|
||||||
@@ -115,9 +122,12 @@ class brushnet_model_loader:
|
|||||||
pbar.update(1)
|
pbar.update(1)
|
||||||
|
|
||||||
#load weights
|
#load weights
|
||||||
brushnet_sd = comfy.utils.load_torch_file(checkpoint_path)
|
brushnet_sd = comfy.utils.load_torch_file(checkpoint_path)
|
||||||
for key in brushnet_sd:
|
if is_accelerate_available():
|
||||||
set_module_tensor_to_device(brushnet, key, device=device, dtype=dtype, value=brushnet_sd[key])
|
for key in brushnet_sd:
|
||||||
|
set_module_tensor_to_device(brushnet, key, device=device, dtype=dtype, value=brushnet_sd[key])
|
||||||
|
else:
|
||||||
|
brushnet.load_state_dict(brushnet_sd)
|
||||||
del brushnet_sd
|
del brushnet_sd
|
||||||
|
|
||||||
clip_sd = None
|
clip_sd = None
|
||||||
@@ -128,14 +138,20 @@ class brushnet_model_loader:
|
|||||||
sd = model.model.state_dict_for_saving(clip_sd, vae.get_sd(), None)
|
sd = model.model.state_dict_for_saving(clip_sd, vae.get_sd(), None)
|
||||||
|
|
||||||
converted_vae = convert_ldm_vae_checkpoint(sd, converted_vae_config)
|
converted_vae = convert_ldm_vae_checkpoint(sd, converted_vae_config)
|
||||||
for key in converted_vae:
|
if is_accelerate_available():
|
||||||
set_module_tensor_to_device(new_vae, key, device=device, dtype=dtype, value=converted_vae[key])
|
for key in converted_vae:
|
||||||
|
set_module_tensor_to_device(new_vae, key, device=device, dtype=dtype, value=converted_vae[key])
|
||||||
|
else:
|
||||||
|
new_vae.load_state_dict(converted_vae)
|
||||||
del converted_vae
|
del converted_vae
|
||||||
pbar.update(1)
|
pbar.update(1)
|
||||||
|
|
||||||
converted_unet = convert_ldm_unet_checkpoint(sd, converted_unet_config)
|
converted_unet = convert_ldm_unet_checkpoint(sd, converted_unet_config)
|
||||||
for key in converted_unet:
|
if is_accelerate_available():
|
||||||
set_module_tensor_to_device(new_unet, key, device=device, dtype=dtype, value=converted_unet[key])
|
for key in converted_unet:
|
||||||
|
set_module_tensor_to_device(new_unet, key, device=device, dtype=dtype, value=converted_unet[key])
|
||||||
|
else:
|
||||||
|
new_unet.load_state_dict(converted_unet)
|
||||||
del converted_unet
|
del converted_unet
|
||||||
|
|
||||||
pbar.update(1)
|
pbar.update(1)
|
||||||
@@ -151,6 +167,7 @@ class brushnet_model_loader:
|
|||||||
|
|
||||||
pbar.update(1)
|
pbar.update(1)
|
||||||
del sd
|
del sd
|
||||||
|
|
||||||
|
|
||||||
self.pipe = StableDiffusionBrushNetPipeline(
|
self.pipe = StableDiffusionBrushNetPipeline(
|
||||||
unet=new_unet,
|
unet=new_unet,
|
||||||
@@ -163,11 +180,18 @@ class brushnet_model_loader:
|
|||||||
safety_checker=None,
|
safety_checker=None,
|
||||||
feature_extractor=None
|
feature_extractor=None
|
||||||
)
|
)
|
||||||
#self.pipe.enable_model_cpu_offload()
|
|
||||||
pbar.update(1)
|
|
||||||
brushnet = {
|
brushnet = {
|
||||||
"pipe": self.pipe,
|
"pipe": self.pipe,
|
||||||
}
|
}
|
||||||
|
if ip_adapter is not None:
|
||||||
|
from .ip_adapter.ip_adapter import IPAdapter
|
||||||
|
brushnet['ip_adapter_weight'] = ip_adapter['ip_adapter_weight']
|
||||||
|
brushnet['ip_adapter_image'] = ip_adapter['ip_adapter_image']
|
||||||
|
|
||||||
|
ip_adapter = IPAdapter(self.pipe, ip_adapter['ipadapter_path'], ip_adapter['image_encoder'], device=device)
|
||||||
|
brushnet['ip_adapter'] = ip_adapter
|
||||||
|
|
||||||
|
pbar.update(1)
|
||||||
|
|
||||||
return (brushnet,)
|
return (brushnet,)
|
||||||
|
|
||||||
@@ -257,19 +281,17 @@ class brushnet_sampler:
|
|||||||
|
|
||||||
B, H, W, C = image.shape
|
B, H, W, C = image.shape
|
||||||
image = image.permute(0, 3, 1, 2).to(device)
|
image = image.permute(0, 3, 1, 2).to(device)
|
||||||
print("mask_shape: ",mask.shape)
|
|
||||||
#handle masks
|
#handle masks
|
||||||
if len(mask.shape) == 2:
|
if len(mask.shape) == 2:
|
||||||
mask = mask.unsqueeze(0)
|
mask = mask.unsqueeze(0)
|
||||||
mask = F.interpolate(mask.unsqueeze(1), size=[H, W], mode='nearest')
|
mask = F.interpolate(mask.unsqueeze(1), size=[H, W], mode='nearest')
|
||||||
mask = mask.to(device)
|
mask = mask.to(device)
|
||||||
print("mask_shape: ",mask.shape)
|
|
||||||
if mask.shape[0] < B:
|
if mask.shape[0] < B:
|
||||||
repeat_times = B // mask.shape[0]
|
repeat_times = B // mask.shape[0]
|
||||||
mask = mask.repeat(repeat_times, 1, 1, 1)
|
mask = mask.repeat(repeat_times, 1, 1, 1)
|
||||||
|
|
||||||
print("mask_shape: ",mask.shape)
|
|
||||||
print("image_shape: ", image.shape)
|
|
||||||
image = image * (1-mask)
|
image = image * (1-mask)
|
||||||
|
|
||||||
if 'ip_adapter' in brushnet:
|
if 'ip_adapter' in brushnet:
|
||||||
@@ -282,7 +304,7 @@ class brushnet_sampler:
|
|||||||
)
|
)
|
||||||
prompt_embeds = torch.repeat_interleave(prompt_embeds, B, dim=0)
|
prompt_embeds = torch.repeat_interleave(prompt_embeds, B, dim=0)
|
||||||
negative_prompt_embeds = torch.repeat_interleave(negative_prompt_embeds, B, dim=0)
|
negative_prompt_embeds = torch.repeat_interleave(negative_prompt_embeds, B, dim=0)
|
||||||
print(prompt_embeds.shape, negative_prompt_embeds.shape)
|
|
||||||
use_ipadapter = True
|
use_ipadapter = True
|
||||||
prompt_list = None
|
prompt_list = None
|
||||||
n_prompt_list = None
|
n_prompt_list = None
|
||||||
@@ -360,7 +382,6 @@ class brushnet_ipadapter_matteo:
|
|||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
return {"required": {
|
return {"required": {
|
||||||
"brushnet": ("BRUSHNET",),
|
|
||||||
"image": ("IMAGE",),
|
"image": ("IMAGE",),
|
||||||
"ipadapter": (folder_paths.get_filename_list("ipadapter"), ),
|
"ipadapter": (folder_paths.get_filename_list("ipadapter"), ),
|
||||||
"clip_vision" : (folder_paths.get_filename_list("clip_vision"), ),
|
"clip_vision" : (folder_paths.get_filename_list("clip_vision"), ),
|
||||||
@@ -368,12 +389,12 @@ class brushnet_ipadapter_matteo:
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("BRUSHNET",)
|
RETURN_TYPES = ("DIFFUSERSIPADAPTER",)
|
||||||
RETURN_NAMES = ("brushnet",)
|
RETURN_NAMES = ("ip_adapter",)
|
||||||
FUNCTION = "loadmodel"
|
FUNCTION = "loadmodel"
|
||||||
CATEGORY = "BrushNetWrapper"
|
CATEGORY = "BrushNetWrapper"
|
||||||
|
|
||||||
def loadmodel(self, image, brushnet, ipadapter, clip_vision, weight):
|
def loadmodel(self, image, ipadapter, clip_vision, weight):
|
||||||
from .ip_adapter.ip_adapter import IPAdapter
|
from .ip_adapter.ip_adapter import IPAdapter
|
||||||
from transformers import CLIPVisionConfig, CLIPVisionModelWithProjection
|
from transformers import CLIPVisionConfig, CLIPVisionModelWithProjection
|
||||||
device = mm.get_torch_device()
|
device = mm.get_torch_device()
|
||||||
@@ -387,17 +408,20 @@ class brushnet_ipadapter_matteo:
|
|||||||
with (init_empty_weights() if is_accelerate_available() else nullcontext()):
|
with (init_empty_weights() if is_accelerate_available() else nullcontext()):
|
||||||
image_encoder = CLIPVisionModelWithProjection(clip_vision_config)
|
image_encoder = CLIPVisionModelWithProjection(clip_vision_config)
|
||||||
clip_vision_sd = comfy.utils.load_torch_file(clip_vision_path)
|
clip_vision_sd = comfy.utils.load_torch_file(clip_vision_path)
|
||||||
for key in clip_vision_sd:
|
if is_accelerate_available():
|
||||||
set_module_tensor_to_device(image_encoder, key, device=device, dtype=dtype, value=clip_vision_sd[key])
|
for key in clip_vision_sd:
|
||||||
|
set_module_tensor_to_device(image_encoder, key, device=device, dtype=dtype, value=clip_vision_sd[key])
|
||||||
brushnet['pipe'].to(device)
|
else:
|
||||||
ip_adapter = IPAdapter(brushnet['pipe'], ipadapter_path, image_encoder, device=device)
|
image_encoder.load_state_dict(clip_vision_sd)
|
||||||
|
|
||||||
|
#ip_adapter = IPAdapter(brushnet['pipe'], ipadapter_path, image_encoder, device=device)
|
||||||
image = image.permute(0, 3, 1, 2).to(device)
|
image = image.permute(0, 3, 1, 2).to(device)
|
||||||
|
ip_adapter = {}
|
||||||
brushnet['ip_adapter'] = ip_adapter
|
ip_adapter['ipadapter_path'] = ipadapter_path
|
||||||
brushnet['ip_adapter_image'] = image
|
ip_adapter['image_encoder'] = image_encoder
|
||||||
brushnet['ip_adapter_weight'] = weight
|
ip_adapter['ip_adapter_image'] = image
|
||||||
return (brushnet,)
|
ip_adapter['ip_adapter_weight'] = weight
|
||||||
|
return (ip_adapter,)
|
||||||
|
|
||||||
class brushnet_sampler_ella:
|
class brushnet_sampler_ella:
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -444,6 +468,8 @@ class brushnet_sampler_ella:
|
|||||||
dtype = mm.unet_dtype()
|
dtype = mm.unet_dtype()
|
||||||
mm.soft_empty_cache()
|
mm.soft_empty_cache()
|
||||||
pipe=brushnet["pipe"].to(dtype)
|
pipe=brushnet["pipe"].to(dtype)
|
||||||
|
if 'ipadapter' in brushnet:
|
||||||
|
raise Exception("This doesn't currently support using both ELLA and IPAdapter.")
|
||||||
|
|
||||||
global IS_MODEL_CPU_OFFLOAD_ENABLED
|
global IS_MODEL_CPU_OFFLOAD_ENABLED
|
||||||
if not IS_MODEL_CPU_OFFLOAD_ENABLED:
|
if not IS_MODEL_CPU_OFFLOAD_ENABLED:
|
||||||
|
|||||||
@@ -1,2 +1,3 @@
|
|||||||
diffusers>=0.27.2
|
diffusers>=0.27.2
|
||||||
|
accelerate>=0.29.2
|
||||||
omegaconf
|
omegaconf
|
||||||
|
|||||||
Reference in New Issue
Block a user