diff --git a/py/cat_vton.py b/py/cat_vton.py index 249e3b4..8b40918 100644 --- a/py/cat_vton.py +++ b/py/cat_vton.py @@ -1,4 +1,5 @@ from .func import * +from comfy.utils import ProgressBar NODE_NAME = 'CatVTON_Wrapper' @@ -69,13 +70,15 @@ class LS_CatVTON: mask = mask_processor.blur(mask, blur_factor=9) # Inference + comfyui_pbar_update = ProgressBar(total=steps).update result_image = pipeline( image=person_image, condition_image=cloth_image, mask=mask, num_inference_steps=steps, guidance_scale=cfg, - generator=generator + generator=generator, + comfy_pbar_callback=comfyui_pbar_update )[0] result_image = restore_padding_image(result_image, target_image.size, person_image_bbox) diff --git a/py/catvton/pipeline.py b/py/catvton/pipeline.py index 4d28015..dcfb5a3 100644 --- a/py/catvton/pipeline.py +++ b/py/catvton/pipeline.py @@ -21,6 +21,7 @@ from .utils1 import ( prepare_mask_image, resize_and_crop, resize_and_padding, + call_callback ) @@ -182,6 +183,10 @@ class CatVTONPipeline: ): progress_bar.update() + comfy_pbar_callback = kwargs.get("comfy_pbar_callback", None) + if comfy_pbar_callback is not None: + call_callback(comfy_pbar_callback, 1) + # Decode the final latents latents = latents.split(latents.shape[concat_dim] // 2, dim=concat_dim)[0] latents = 1 / self.vae.config.scaling_factor * latents diff --git a/py/catvton/utils.py b/py/catvton/utils.py index 9cbb344..3fef361 100644 --- a/py/catvton/utils.py +++ b/py/catvton/utils.py @@ -79,6 +79,3 @@ def get_trainable_module(unet, trainable_module_name): return attn_blocks else: raise ValueError(f"Unknown trainable_module_name: {trainable_module_name}") - - - diff --git a/py/catvton/utils1.py b/py/catvton/utils1.py index 9e19809..fb88fef 100644 --- a/py/catvton/utils1.py +++ b/py/catvton/utils1.py @@ -667,3 +667,6 @@ if __name__ == "__main__": ) vis_sobel_weight(image_path, mask_path).save(result_path) pass + +def call_callback(callback:callable, *args, **kwargs): + return callback(*args, **kwargs)