init
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
@@ -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**
|
||||
> <br>
|
||||
> [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)
|
||||
> <br>
|
||||
> [Show Lab](https://sites.google.com/view/showlab), National University of Singapore
|
||||
> <br>
|
||||
|
||||
# Example
|
||||

|
||||
<a href="https://arxiv.org/abs/2502.14397"><img src="https://img.shields.io/badge/ariXv-2502.14397-A42C25.svg" alt="arXiv"></a>
|
||||
<a href="https://huggingface.co/nicolaus-huang/PhotoDoodle"><img src="https://img.shields.io/badge/🤗_HuggingFace-Model-ffbd45.svg" alt="HuggingFace"></a>
|
||||
<a href="https://huggingface.co/datasets/nicolaus-huang/PhotoDoodle/"><img src="https://img.shields.io/badge/🤗_HuggingFace-Dataset-ffbd45.svg" alt="HuggingFace"></a>
|
||||
|
||||
# Citation
|
||||
<br>
|
||||
|
||||
<img src='./assets/teaser.png' width='100%' />
|
||||
|
||||
|
||||
## 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
|
||||
<span id="dataset_setting"></span>
|
||||
#### 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},
|
||||
}
|
||||
``
|
||||
```
|
||||
|
||||
@@ -0,0 +1,4 @@
|
||||
|
||||
from .PhotoDoodle_node import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
|
||||
@@ -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()
|
||||
@@ -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")
|
||||
+213
@@ -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
|
||||
@@ -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 = ""
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user