From b1a35ba4a88a05c5938831d278231b0d979e1fe9 Mon Sep 17 00:00:00 2001 From: camd11 Date: Mon, 5 May 2025 17:32:45 -0400 Subject: [PATCH] fix cpu offloading --- InstantCharacter/pipeline.py | 14 ++++++++++---- nodes/comfy_nodes.py | 35 ++++++++++++++++++++++------------- 2 files changed, 32 insertions(+), 17 deletions(-) diff --git a/InstantCharacter/pipeline.py b/InstantCharacter/pipeline.py index 1ea1082..2681c73 100644 --- a/InstantCharacter/pipeline.py +++ b/InstantCharacter/pipeline.py @@ -35,6 +35,8 @@ class InstantCharacterFluxPipeline(FluxPipeline): @torch.inference_mode() def encode_siglip_image_emb(self, siglip_image, device, dtype): + # Ensure encoder is on the correct device before use + self.siglip_image_encoder.to(device, dtype=dtype) siglip_image = siglip_image.to(device, dtype=dtype) res = self.siglip_image_encoder(siglip_image, output_hidden_states=True) @@ -47,6 +49,8 @@ class InstantCharacterFluxPipeline(FluxPipeline): @torch.inference_mode() def encode_dinov2_image_emb(self, dinov2_image, device, dtype): + # Ensure encoder is on the correct device before use + self.dino_image_encoder_2.to(device, dtype=dtype) dinov2_image = dinov2_image.to(device, dtype=dtype) res = self.dino_image_encoder_2(dinov2_image, output_hidden_states=True) @@ -457,11 +461,13 @@ class InstantCharacterFluxPipeline(FluxPipeline): # subject adapter if subject_image is not None: + # Ensure projector is on the correct device before use when offloading + self.subject_image_proj_model.to(latents.device, dtype=latents.dtype) subject_image_prompt_embeds = self.subject_image_proj_model( - low_res_shallow=subject_image_embeds_dict['image_embeds_low_res_shallow'], - low_res_deep=subject_image_embeds_dict['image_embeds_low_res_deep'], - high_res_deep=subject_image_embeds_dict['image_embeds_high_res_deep'], - timesteps=timestep.to(dtype=latents.dtype), + low_res_shallow=subject_image_embeds_dict['image_embeds_low_res_shallow'].to(latents.device, dtype=latents.dtype), + low_res_deep=subject_image_embeds_dict['image_embeds_low_res_deep'].to(latents.device, dtype=latents.dtype), + high_res_deep=subject_image_embeds_dict['image_embeds_high_res_deep'].to(latents.device, dtype=latents.dtype), + timesteps=timestep.to(device=latents.device, dtype=latents.dtype), need_temb=True )[0] self._joint_attention_kwargs['emb_dict'] = dict( diff --git a/nodes/comfy_nodes.py b/nodes/comfy_nodes.py index 4a32af3..d86b391 100644 --- a/nodes/comfy_nodes.py +++ b/nodes/comfy_nodes.py @@ -45,16 +45,20 @@ class InstantCharacterLoadModelFromLocal: pipe = InstantCharacterFluxPipeline.from_pretrained(base_model_path, torch_dtype=torch.bfloat16) + # Initialize adapter first + pipe.init_adapter( + image_encoder_path=image_encoder_path, + image_encoder_2_path=image_encoder_2_path, + subject_ipadapter_cfg=dict(subject_ip_adapter_path=ip_adapter_path, nb_token=1024), + ) + + # Then move to device or enable offloading if cpu_offload: + print("Enabling CPU offload for InstantCharacter pipeline...") pipe.enable_sequential_cpu_offload() + print("CPU offload enabled.") else: pipe.to(device) - - pipe.init_adapter( - image_encoder_path=image_encoder_path, - image_encoder_2_path=image_encoder_2_path, - subject_ipadapter_cfg=dict(subject_ip_adapter_path=ip_adapter_path, nb_token=1024), - ) return (pipe,) @@ -88,13 +92,10 @@ class InstantCharacterLoadModel: pipe = InstantCharacterFluxPipeline.from_pretrained( base_model, torch_dtype=torch.bfloat16, - cache_dir=cache_dir, + cache_dir=cache_dir, ) - if cpu_offload: - pipe.enable_sequential_cpu_offload() - else: - pipe.to(device) - + + # Initialize adapter first pipe.init_adapter( image_encoder_path=image_encoder_path, cache_dir=image_encoder_cache_dir, @@ -105,7 +106,15 @@ class InstantCharacterLoadModel: nb_token=1024 ), ) - + + # Then move to device or enable offloading + if cpu_offload: + print("Enabling CPU offload for InstantCharacter pipeline...") + pipe.enable_sequential_cpu_offload() + print("CPU offload enabled.") + else: + pipe.to(device) + return (pipe,)