Merge pull request #14 from camd11/main

fix cpu offloading
This commit is contained in:
jax
2025-05-11 09:49:11 +08:00
committed by GitHub
2 changed files with 32 additions and 17 deletions
+10 -4
View File
@@ -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(
+22 -13
View File
@@ -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,)