From 9852178282b09e3e80ff6a648bb44c2ed025f8e4 Mon Sep 17 00:00:00 2001
From: smthemex <138738845+smthemex@users.noreply.github.com>
Date: Thu, 27 Feb 2025 19:17:18 +0800
Subject: [PATCH] init
---
LICENSE | 2 +-
PhotoDoodle_node.py | 191 +++++++++++++++++++++++++++++++++++++++
README.md | 134 ++++++++++++++++++++++++++--
__init__.py | 4 +
inference.py | 69 ++++++++++++++
merge.py | 13 +++
node_utils.py | 213 ++++++++++++++++++++++++++++++++++++++++++++
pyproject.toml | 15 ++++
requirements.txt | 29 ++++++
9 files changed, 660 insertions(+), 10 deletions(-)
create mode 100644 PhotoDoodle_node.py
create mode 100644 __init__.py
create mode 100644 inference.py
create mode 100644 merge.py
create mode 100644 node_utils.py
create mode 100644 pyproject.toml
create mode 100644 requirements.txt
diff --git a/LICENSE b/LICENSE
index a68732e..dd76e79 100644
--- a/LICENSE
+++ b/LICENSE
@@ -1,6 +1,6 @@
MIT License
-Copyright (c) 2025 smthemex
+Copyright (c) 2025 Show Lab
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
diff --git a/PhotoDoodle_node.py b/PhotoDoodle_node.py
new file mode 100644
index 0000000..5be8285
--- /dev/null
+++ b/PhotoDoodle_node.py
@@ -0,0 +1,191 @@
+# !/usr/bin/env python
+# -*- coding: UTF-8 -*-
+import os
+import torch
+import numpy as np
+from diffusers import AutoencoderKL,FluxTransformer2DModel
+
+from .node_utils import tensor2pil_list, load_images,cleanup
+from .src.pipeline_pe_clone import FluxPipeline
+from .src.pipeline_pe_clone_orgin import FluxPipeline as FluxPipeline_orgin
+
+import folder_paths
+from comfy.utils import ProgressBar
+MAX_SEED = np.iinfo(np.int32).max
+current_node_path = os.path.dirname(os.path.abspath(__file__))
+
+device = torch.device(
+ "cuda:0") if torch.cuda.is_available() else torch.device(
+ "mps") if torch.backends.mps.is_available() else torch.device(
+ "cpu")
+
+
+class PhotoDoodle_Loader:
+ def __init__(self):
+ pass
+
+ @classmethod
+ def INPUT_TYPES(s):
+ return {
+ "required": {
+ "flux_unet": (["none"] + folder_paths.get_filename_list("diffusion_models"),),
+ "vae": (["none"] + folder_paths.get_filename_list("vae"),),
+ "pre_lora": (["none"] + [i for i in folder_paths.get_filename_list("loras") if "pre" in i],),
+ "loras": (["none"] + folder_paths.get_filename_list("loras"),),
+ "flux_repo":("STRING", {"default": "", "multiline": False}),
+ "use_mmgp":("BOOLEAN",{"default":False}),
+ "profile_number":([0,1,2,3,4,5],)
+ },
+ }
+
+ RETURN_TYPES = ("MODEL_PhotoDoodle",)
+ RETURN_NAMES = ("model",)
+ FUNCTION = "loader_main"
+ CATEGORY = "PhotoDoodle"
+
+
+ def loader_main(self,flux_unet,vae,pre_lora,loras,flux_repo,use_mmgp,profile_number):
+
+ flux_repo_local=os.path.join(current_node_path, 'src/FLUX.1-dev')
+ print("***********Load model ***********")
+ #load model
+
+ if pre_lora != "none":
+ pre_lora_path = folder_paths.get_full_path("loras", pre_lora)
+ else:
+ raise ValueError("No model selected")
+
+ if flux_unet != "none":
+ flux_transformer_path = folder_paths.get_full_path("diffusion_models", flux_unet)
+ else:
+ raise ValueError("No model selected")
+
+ if loras != "none":
+ lora_path = folder_paths.get_full_path("loras", loras)
+ else:
+ raise ValueError("No model selected")
+ need_clip=False
+ if not flux_repo:
+ if vae != "none":
+ vae_path = folder_paths.get_full_path("vae", vae)
+ vae_config=os.path.join(flux_repo_local, 'vae')
+ ae = AutoencoderKL.from_single_file(vae_path,config=vae_config, torch_dtype=torch.bfloat16)
+ transformer = FluxTransformer2DModel.from_single_file(
+ flux_transformer_path,
+ config=os.path.join(flux_repo_local, "transformer"),
+ torch_dtype=torch.bfloat16,
+ )
+ pipeline = FluxPipeline.from_pretrained(
+ flux_repo_local,
+ vae=ae,
+ transformer=transformer,
+ torch_dtype=torch.bfloat16,
+ )
+ need_clip=True
+ else:
+ pipeline = FluxPipeline_orgin.from_single_file(
+ flux_transformer_path,
+ config=flux_repo_local,
+ torch_dtype=torch.bfloat16,
+ )
+
+ else:
+ # flux_repo='F:/test/ComfyUI/models/diffusers/black-forest-labs/FLUX.1-dev'
+ pipeline =FluxPipeline_orgin.from_pretrained(flux_repo,torch_dtype=torch.bfloat16,)
+
+ # Load and fuse base LoRA weights
+ pipeline.load_lora_weights(lora_path)
+ pipeline.fuse_lora()
+ pipeline.unload_lora_weights()
+ pipeline.load_lora_weights(pre_lora_path)
+ pipeline.enable_model_cpu_offload()
+
+ print("***********Load model done ***********")
+
+ cleanup()
+
+ if use_mmgp:
+ from mmgp import offload as offloadobj
+ pipe = {"transformer": pipeline, }
+ # offloadobj.profile(pipe, quantizeTransformer = False, profile_no = 1 ) # uncomment this line and comment the previous one if you have 24 GB of VRAM and wants faster generation
+ offloadobj.profile(pipe, quantizeTransformer = False, extraModelsToQuantize = [], profile_no = profile_number, )
+ return ({"pipeline":pipeline,"need_clip":need_clip},)
+
+
+
+class PhotoDoodle_Sampler:
+ def __init__(self):
+ pass
+
+ @classmethod
+ def INPUT_TYPES(s):
+ return {
+ "required": {
+ "model": ("MODEL_PhotoDoodle",),
+ "images": ("IMAGE",),
+ "prompt": ("STRING", {"default": "add a halo and wings for the cat by sksmagiceffects", "multiline": True}),
+ "width": ("INT", {"default": 512, "min": 256, "max": 4096, "step": 64, "display": "number"}),
+ "height": ("INT", {"default": 768, "min": 256, "max": 4096, "step": 64, "display": "number"}),
+ "steps": ("INT", {"default": 20, "min": 1, "max": 1024, "step": 1, "display": "number"}),
+ "guidance_scale": ("FLOAT", {"default": 3.5, "min": 0.0, "max": 10.0, "step": 0.1}),
+ "max_sequence_length": ("INT", {"default": 512, "min": 128, "max": 512, "step": 1, "display": "number"}),
+ },
+ "optional": { "clip":("CLIP",),},
+ }
+
+ RETURN_TYPES = ("IMAGE", )
+ RETURN_NAMES = ("image",)
+ FUNCTION = "sampler_main"
+ CATEGORY = "PhotoDoodle"
+
+ def sampler_main(self, model,images,prompt, width, height,steps,guidance_scale, max_sequence_length,**kwargs):
+
+ need_clip=model.get("need_clip")
+ pipeline=model.get("pipeline")
+ clip=kwargs.get("clip")
+ # load model
+ #model.to(device)
+ if need_clip:
+ if clip is None:
+ raise ValueError("No clip selected")
+ prompt_embeds,pooled_prompt_embeds=clip.encode_from_tokens( clip.tokenize(prompt,return_word_ids=True),return_pooled=True)
+ prompt_embeds=prompt_embeds.to(device,torch.bfloat16)
+ pooled_prompt_embeds=pooled_prompt_embeds.to(device,torch.bfloat16)
+ prompt=None
+ else:
+ prompt_embeds,pooled_prompt_embeds=None,None
+
+ cleanup()
+ condition_image_list = tensor2pil_list(images, width, height)
+ total_images=len(condition_image_list)
+ pbar = ProgressBar(total_images)
+ img_list=[]
+
+ for i, condition_image in enumerate(condition_image_list):
+ result = pipeline(
+ prompt=prompt,
+ prompt_embeds=prompt_embeds,
+ pooled_prompt_embeds=pooled_prompt_embeds,
+ condition_image=condition_image,
+ height=height,
+ width=width,
+ guidance_scale=guidance_scale,
+ num_inference_steps=steps,
+ max_sequence_length=max_sequence_length,
+ ).images[0]
+ pbar.update_absolute(i, total_images)
+ img_list.append(result)
+
+ cleanup()
+ return (load_images(img_list),)
+
+
+NODE_CLASS_MAPPINGS = {
+ "PhotoDoodle_Loader": PhotoDoodle_Loader,
+ "PhotoDoodle_Sampler": PhotoDoodle_Sampler,
+}
+
+NODE_DISPLAY_NAME_MAPPINGS = {
+ "PhotoDoodle_Loader": "PhotoDoodle_Loader",
+ "PhotoDoodle_Sampler": "PhotoDoodle_Sampler",
+}
diff --git a/README.md b/README.md
index af5f4cd..8247024 100644
--- a/README.md
+++ b/README.md
@@ -1,14 +1,130 @@
-# ComfyUI_PhotoDoodle
-[PhotoDoodle](https://github.com/showlab/PhotoDoodle): Learning Artistic Image Editing from Few-Shot Pairwise Data,you can use it in comfyUI
+# PhotoDoodle
-# Tips
-* Need Vram>=12G(origin repositories need 40G Vram),coming soon...
-* 12G显存起步,8G未测试,毕竟官方要求40G起步,实在不好优化,晚上上线,白天忙。
+> **PhotoDoodle: Learning Artistic Image Editing from Few-Shot Pairwise Data**
+>
+> [Huang Shijie](https://scholar.google.com/citations?user=HmqYYosAAAAJ),
+> [Yiren Song](https://scholar.google.com.hk/citations?user=L2YS0jgAAAAJ),
+> [Yuxuan Zhang](https://xiaojiu-z.github.io/YuxuanZhang.github.io/),
+> [Hailong Guo](https://github.com/logn-2024),
+> Xueyin Wang,
+> and
+> [Mike Zheng Shou](https://sites.google.com/view/showlab),
+> [Liu Jiaming](https://scholar.google.com/citations?user=SmL7oMQAAAAJ&hl=en)
+>
+> [Show Lab](https://sites.google.com/view/showlab), National University of Singapore
+>
-# Example
-
+
+
+
-# Citation
+
+
+
+
+
+## Quick Start
+### Configuration
+#### 1. **Environment setup**
+```bash
+git clone git@github.com:showlab/PhotoDoodle.git
+cd PhotoDoodle
+
+conda create -n doodle python=3.11.10
+conda activate doodle
+```
+#### 2. **Requirements installation**
+```bash
+pip install torch==2.5.1 torchvision==0.20.1 torchaudio==2.5.1 --index-url https://download.pytorch.org/whl/cu124
+pip install --upgrade -r requirements.txt
+```
+
+
+### 2. Inference
+We provided the intergration of diffusers pipeline with our model and uploaded the model weights to huggingface, it's easy to use the our model as example below:
+
+```bash
+from src.pipeline_pe_clone import FluxPipeline
+import torch
+from PIL import Image
+
+pretrained_model_name_or_path = "black-forest-labs/FLUX.1-dev"
+pipeline = FluxPipeline.from_pretrained(
+ pretrained_model_name_or_path,
+ torch_dtype=torch.bfloat16,
+).to('cuda')
+
+pipeline.load_lora_weights("nicolaus-huang/PhotoDoodle", weight_name="pretrain.safetensors")
+pipeline.fuse_lora()
+pipeline.unload_lora_weights()
+
+pipeline.load_lora_weights("nicolaus-huang/PhotoDoodle", weight_name="sksmagiceffects.safetensors")
+
+height=768
+width=512
+
+validation_image = "assets/1.png"
+validation_prompt = "add a halo and wings for the cat by sksmagiceffects"
+condition_image = Image.open(validation_image).resize((height, width)).convert("RGB")
+
+result = pipeline(prompt=validation_prompt,
+ condition_image=condition_image,
+ height=height,
+ width=width,
+ guidance_scale=3.5,
+ num_inference_steps=20,
+ max_sequence_length=512).images[0]
+
+result.save("output.png")
+```
+
+or simply run the inference script:
+```
+python inference.py
+```
+
+
+
+### 3. Weights
+You can download the trained checkpoints of PhotoDoodle for inference. Below are the details of available models, checkpoint name are also trigger words.
+
+You would need to load and fuse the `pretrained ` checkpoints model in order to load the other models.
+
+| **Model** | **Description** | **Resolution** |
+| :----------------------------------------------------------: | :---------------------------------------------------------: | :------------: |
+| [pretrained](https://huggingface.co/nicolaus-huang/PhotoDoodle/blob/main/pretrain.safetensors) | PhotoDoodle model trained on `SeedEdit` dataset | 768, 768 |
+| [sksmonstercalledlulu](https://huggingface.co/nicolaus-huang/PhotoDoodle/blob/main/sksmonstercalledlulu.safetensors) | PhotoDoodle model trained on `Cartoon monster` dataset | 768, 512 |
+| [sksmagiceffects](https://huggingface.co/nicolaus-huang/PhotoDoodle/blob/main/sksmagiceffects.safetensors) | PhotoDoodle model trained on `3D effects` dataset | 768, 512 |
+| [skspaintingeffects ](https://huggingface.co/nicolaus-huang/PhotoDoodle/blob/main/skspaintingeffects.safetensors) | PhotoDoodle model trained on `Flowing color blocks` dataset | 768, 512 |
+| [sksedgeeffect ](https://huggingface.co/nicolaus-huang/PhotoDoodle/blob/main/sksedgeeffect.safetensors) | PhotoDoodle model trained on `Hand-drawn outline` dataset | 768, 512 |
+
+
+### 4. Dataset
+
+#### 2.1 Settings for dataset
+The training process uses a paired dataset stored in a .jsonl file, where each entry contains image file paths and corresponding text descriptions. Each entry includes the source image path, the target (modified) image path, and a caption describing the modification.
+
+Example format:
+
+```json
+{"source": "path/to/source.jpg", "target": "path/to/modified.jpg", "caption": "Instruction of modifications"}
+{"source": "path/to/source2.jpg", "target": "path/to/modified2.jpg", "caption": "Another instruction"}
+```
+
+We have uploaded our datasets to [Hugging Face](https://huggingface.co/datasets/nicolaus-huang/PhotoDoodle).
+
+
+### 5. Results
+
+
+
+
+### 6. Acknowledgments
+
+1. Thanks to **[Yuxuan Zhang](https://xiaojiu-z.github.io/YuxuanZhang.github.io/)** and **[Hailong Guo](mailto:guohailong@bupt.edu.cn)** for providing the code base.
+2. Thanks to **[Diffusers](https://github.com/huggingface/diffusers)** for the open-source project.
+
+## Citation
```
@misc{huang2025photodoodlelearningartisticimage,
title={PhotoDoodle: Learning Artistic Image Editing from Few-Shot Pairwise Data},
@@ -19,4 +135,4 @@
primaryClass={cs.CV},
url={https://arxiv.org/abs/2502.14397},
}
-``
+```
diff --git a/__init__.py b/__init__.py
new file mode 100644
index 0000000..2264f1e
--- /dev/null
+++ b/__init__.py
@@ -0,0 +1,4 @@
+
+from .PhotoDoodle_node import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
+
+__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
diff --git a/inference.py b/inference.py
new file mode 100644
index 0000000..95313d6
--- /dev/null
+++ b/inference.py
@@ -0,0 +1,69 @@
+import argparse
+from .src.pipeline_pe_clone import FluxPipeline
+import torch
+from PIL import Image
+
+def parse_args():
+ parser = argparse.ArgumentParser(description='FLUX image generation with LoRA')
+ parser.add_argument('--model_path', type=str,
+ default="black-forest-labs/FLUX.1-dev",
+ help='Path to pretrained model')
+ parser.add_argument('--image_path', type=str,
+ default="assets/1.png",
+ help='Input image path')
+ parser.add_argument('--output_path', type=str,
+ default="output.png",
+ help='Output image path')
+ parser.add_argument('--height', type=int, default=768)
+ parser.add_argument('--width', type=int, default=512)
+ parser.add_argument('--prompt', type=str,
+ default="add a halo and wings for the cat by sksmagiceffects",
+ help="""Different LoRA effects and their example prompts:
+ - sksmagiceffects: "add a halo and wings for the cat by sksmagiceffects"
+ - sksmonstercalledlulu: "add a red sksmonstercalledlulu hugging the cat"
+ - skspaintingeffects: "add a yellow flower on the cat's head and psychedelic colors and dynamic flows by skspaintingeffects"
+ - sksedgeeffect: "add yellow flames to the cat by sksedgeeffect"
+ """)
+ parser.add_argument('--guidance_scale', type=float, default=3.5)
+ parser.add_argument('--num_steps', type=int, default=20,
+ help='Number of inference steps')
+ parser.add_argument('--lora_name', type=str,
+ choices=['pretrained', 'sksmagiceffects', 'sksmonstercalledlulu',
+ 'skspaintingeffects', 'sksedgeeffect'],
+ default="sksmagiceffects",
+ help='Name of LoRA weights to use. Use "pretrained" for base model only')
+ return parser.parse_args()
+
+def main():
+ args = parse_args()
+
+ pipeline = FluxPipeline.from_pretrained(
+ args.model_path,
+ torch_dtype=torch.bfloat16,
+ ).to('cuda')
+
+ # Load and fuse base LoRA weights
+ pipeline.load_lora_weights("nicolaus-huang/PhotoDoodle", weight_name="pretrain.safetensors")
+ pipeline.fuse_lora()
+ pipeline.unload_lora_weights()
+
+ # Load selected LoRA effect only if not using pretrained
+ if args.lora_name != 'pretrained':
+ pipeline.load_lora_weights("nicolaus-huang/PhotoDoodle", weight_name=f"{args.lora_name}.safetensors")
+
+ condition_image = Image.open(args.image_path).resize((args.height, args.width)).convert("RGB")
+
+ result = pipeline(
+ prompt=args.prompt,
+ condition_image=condition_image,
+ height=args.height,
+ width=args.width,
+ guidance_scale=args.guidance_scale,
+ num_inference_steps=args.num_steps,
+ max_sequence_length=512
+ ).images[0]
+
+ result.save(args.output_path)
+
+if __name__ == "__main__":
+ main()
diff --git a/merge.py b/merge.py
new file mode 100644
index 0000000..e875d78
--- /dev/null
+++ b/merge.py
@@ -0,0 +1,13 @@
+
+from .src.pipeline_pe_clone import FluxPipeline
+import torch
+from PIL import Image
+pretrained_model_name_or_path = "black-forest-labs/FLUX.1-dev"
+pipeline = FluxPipeline.from_pretrained(
+ pretrained_model_name_or_path,
+ torch_dtype=torch.bfloat16,
+)
+pipeline.load_lora_weights("outputs/doodle_pretrain_4508000/pytorch_lora_weights.safetensors")
+pipeline.fuse_lora()
+pipeline.unload_lora_weights()
+pipeline.save_pretrained("edit_pretrain")
\ No newline at end of file
diff --git a/node_utils.py b/node_utils.py
new file mode 100644
index 0000000..78a80af
--- /dev/null
+++ b/node_utils.py
@@ -0,0 +1,213 @@
+# !/usr/bin/env python
+# -*- coding: UTF-8 -*-
+import os
+import torch
+from PIL import Image
+import numpy as np
+import cv2
+import gc
+
+from comfy.utils import common_upscale,ProgressBar
+from huggingface_hub import hf_hub_download
+
+cur_path = os.path.dirname(os.path.abspath(__file__))
+device = "cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu"
+
+
+def cleanup():
+ gc.collect()
+ torch.cuda.empty_cache()
+
+def cv2pil(cv_image):
+ """
+ 将OpenCV图像转换为PIL图像
+ :param cv_image: OpenCV图像
+ :return: PIL图像
+ """
+ # 将图像从BGR转换为RGB
+ rgb_image = cv2.cvtColor(cv_image, cv2.COLOR_BGR2RGB)
+ # 使用PIL的Image.fromarray方法将NumPy数组转换为PIL图像
+ pil_image = Image.fromarray(rgb_image)
+ return pil_image
+
+
+def tensor_to_pil(tensor):
+ image_np = tensor.squeeze().mul(255).clamp(0, 255).byte().numpy()
+ image = Image.fromarray(image_np, mode='RGB')
+ return image
+
+def tensor2pil_list(image,width,height):
+ B,_,_,_=image.size()
+ if B==1:
+ ref_image_list=[tensor2pil_upscale(image,width,height)]
+ else:
+ img_list = list(torch.chunk(image, chunks=B))
+ ref_image_list = [tensor2pil_upscale(img,width,height) for img in img_list]
+ return ref_image_list
+
+
+def tensor_upscale(img_tensor, width, height):
+ samples = img_tensor.movedim(-1, 1)
+ img = common_upscale(samples, width, height, "nearest-exact", "center")
+ samples = img.movedim(1, -1)
+ return samples
+
+def tensor2pil_upscale(img_tensor, width, height):
+ samples = img_tensor.movedim(-1, 1)
+ img = common_upscale(samples, width, height, "nearest-exact", "center")
+ samples = img.movedim(1, -1)
+ img_pil = tensor_to_pil(samples)
+ return img_pil
+
+
+def tensor2cv(tensor_image,RGB2BGR=True):
+ if len(tensor_image.shape)==4:#bhwc to hwc
+ tensor_image=tensor_image.squeeze(0)
+ if tensor_image.is_cuda:
+ tensor_image = tensor_image.cpu().detach()
+ tensor_image=tensor_image.numpy()
+ #反归一化
+ maxValue=tensor_image.max()
+ tensor_image=tensor_image*255/maxValue
+ img_cv2=np.uint8(tensor_image)#32 to uint8
+ if RGB2BGR:
+ img_cv2=cv2.cvtColor(img_cv2,cv2.COLOR_RGB2BGR)
+ return img_cv2
+
+def cvargb2tensor(img):
+ assert type(img) == np.ndarray, 'the img type is {}, but ndarry expected'.format(type(img))
+ img = torch.from_numpy(img.transpose((2, 0, 1)))
+ return img.float().div(255).unsqueeze(0) # 255也可以改为256
+
+def cv2tensor(img):
+ assert type(img) == np.ndarray, 'the img type is {}, but ndarry expected'.format(type(img))
+ img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
+ img = torch.from_numpy(img.transpose((2, 0, 1)))
+ return img.float().div(255).unsqueeze(0) # 255也可以改为256
+
+def images_generator(img_list: list,):
+ #get img size
+ sizes = {}
+ for image_ in img_list:
+ if isinstance(image_,Image.Image):
+ count = sizes.get(image_.size, 0)
+ sizes[image_.size] = count + 1
+ elif isinstance(image_,np.ndarray):
+ count = sizes.get(image_.shape[:2][::-1], 0)
+ sizes[image_.shape[:2][::-1]] = count + 1
+ else:
+ raise "unsupport image list,must be pil or cv2!!!"
+ size = max(sizes.items(), key=lambda x: x[1])[0]
+ yield size[0], size[1]
+
+ # any to tensor
+ def load_image(img_in):
+ if isinstance(img_in, Image.Image):
+ img_in=img_in.convert("RGB")
+ i = np.array(img_in, dtype=np.float32)
+ i = torch.from_numpy(i).div_(255)
+ if i.shape[0] != size[1] or i.shape[1] != size[0]:
+ i = torch.from_numpy(i).movedim(-1, 0).unsqueeze(0)
+ i = common_upscale(i, size[0], size[1], "lanczos", "center")
+ i = i.squeeze(0).movedim(0, -1).numpy()
+ return i
+ elif isinstance(img_in,np.ndarray):
+ i=cv2.cvtColor(img_in,cv2.COLOR_BGR2RGB).astype(np.float32)
+ i = torch.from_numpy(i).div_(255)
+ #print(i.shape)
+ return i
+ else:
+ raise "unsupport image list,must be pil,cv2 or tensor!!!"
+
+ total_images = len(img_list)
+ processed_images = 0
+ pbar = ProgressBar(total_images)
+ images = map(load_image, img_list)
+ try:
+ prev_image = next(images)
+ while True:
+ next_image = next(images)
+ yield prev_image
+ processed_images += 1
+ pbar.update_absolute(processed_images, total_images)
+ prev_image = next_image
+ except StopIteration:
+ pass
+ if prev_image is not None:
+ yield prev_image
+
+def load_images(img_list: list,):
+ gen = images_generator(img_list)
+ (width, height) = next(gen)
+ images = torch.from_numpy(np.fromiter(gen, np.dtype((np.float32, (height, width, 3)))))
+ if len(images) == 0:
+ raise FileNotFoundError(f"No images could be loaded .")
+ return images
+
+def tensor2pil(tensor):
+ image_np = tensor.squeeze().mul(255).clamp(0, 255).byte().numpy()
+ image = Image.fromarray(image_np, mode='RGB')
+ return image
+
+def pil2narry(img):
+ narry = torch.from_numpy(np.array(img).astype(np.float32) / 255.0).unsqueeze(0)
+ return narry
+
+def equalize_lists(list1, list2):
+ """
+ 比较两个列表的长度,如果不一致,则将较短的列表复制以匹配较长列表的长度。
+
+ 参数:
+ list1 (list): 第一个列表
+ list2 (list): 第二个列表
+
+ 返回:
+ tuple: 包含两个长度相等的列表的元组
+ """
+ len1 = len(list1)
+ len2 = len(list2)
+
+ if len1 == len2:
+ pass
+ elif len1 < len2:
+ print("list1 is shorter than list2, copying list1 to match list2's length.")
+ list1.extend(list1 * ((len2 // len1) + 1)) # 复制list1以匹配list2的长度
+ list1 = list1[:len2] # 确保长度一致
+ else:
+ print("list2 is shorter than list1, copying list2 to match list1's length.")
+ list2.extend(list2 * ((len1 // len2) + 1)) # 复制list2以匹配list1的长度
+ list2 = list2[:len1] # 确保长度一致
+
+ return list1, list2
+
+def file_exists(directory, filename):
+ # 构建文件的完整路径
+ file_path = os.path.join(directory, filename)
+ # 检查文件是否存在
+ return os.path.isfile(file_path)
+
+def download_weights(file_dir,repo_id,subfolder="",pt_name=""):
+ if subfolder:
+ file_path = os.path.join(file_dir,subfolder, pt_name)
+ sub_dir=os.path.join(file_dir,subfolder)
+ if not os.path.exists(sub_dir):
+ os.makedirs(sub_dir)
+ if not os.path.exists(file_path):
+ file_path = hf_hub_download(
+ repo_id=repo_id,
+ subfolder=subfolder,
+ filename=pt_name,
+ local_dir = file_dir,
+ )
+ return file_path
+ else:
+ file_path = os.path.join(file_dir, pt_name)
+ if not os.path.exists(file_dir):
+ os.makedirs(file_dir)
+ if not os.path.exists(file_path):
+ file_path = hf_hub_download(
+ repo_id=repo_id,
+ filename=pt_name,
+ local_dir=file_dir,
+ )
+ return file_path
diff --git a/pyproject.toml b/pyproject.toml
new file mode 100644
index 0000000..e1df07c
--- /dev/null
+++ b/pyproject.toml
@@ -0,0 +1,15 @@
+[project]
+name = "comfyui_photodoodel"
+description = "PhotoDoodle: Learning Artistic Image Editing from Few-Shot Pairwise Data,you can use it in comfyUI"
+version = "1.0.0"
+license = {file = "LICENSE"}
+dependencies = ["accelerate>=0.33.0", "transformers>=4.44.0", "diffusers>=0.31.0", "#diffusers[torch]==0.25.0", "#ftfy==6.1.1", "# albumentations==1.3.0", "#opencv-python==4.8.1.78", "#einops==0.7.0", "#pytorch-lightning==1.9.0", "bitsandbytes>=0.44.0", "#prodigyopt==1.0", "#lion-pytorch==0.0.6", "#came_pytorch==0.1.3", "#schedulefree==1.4", "#tensorboard", "#safetensors==0.4.4", "# for gradio", "#gradio==3.6", "#altair==4.2.2", "#easygui==0.98.3", "#toml==0.10.2", "#voluptuous==0.13.1", "#huggingface-hub==0.24.5", "# for Image utils", "#imagesize==1.4.1", "#numpy<=2.0", "#rich==13.7.0", "# for T5XXL tokenizer (SD3/FLUX)", "#sentencepiece==0.2.0"]
+
+[project.urls]
+Repository = "https://github.com/smthemex/ComfyUI_PhotoDoodle"
+# Used by Comfy Registry https://comfyregistry.org
+
+[tool.comfy]
+PublisherId = "smthemex"
+DisplayName = "ComfyUI_PhotoDoodle"
+Icon = ""
diff --git a/requirements.txt b/requirements.txt
new file mode 100644
index 0000000..4229050
--- /dev/null
+++ b/requirements.txt
@@ -0,0 +1,29 @@
+accelerate>=0.33.0
+transformers>=4.44.0
+diffusers>=0.31.0
+#diffusers[torch]==0.25.0
+#ftfy==6.1.1
+# albumentations==1.3.0
+#opencv-python==4.8.1.78
+#einops==0.7.0
+#pytorch-lightning==1.9.0
+bitsandbytes>=0.44.0
+#prodigyopt==1.0
+#lion-pytorch==0.0.6
+#came_pytorch==0.1.3
+#schedulefree==1.4
+#tensorboard
+#safetensors==0.4.4
+# for gradio
+#gradio==3.6
+#altair==4.2.2
+#easygui==0.98.3
+#toml==0.10.2
+#voluptuous==0.13.1
+#huggingface-hub==0.24.5
+# for Image utils
+#imagesize==1.4.1
+#numpy<=2.0
+#rich==13.7.0
+# for T5XXL tokenizer (SD3/FLUX)
+#sentencepiece==0.2.0
\ No newline at end of file