Files
2025-05-21 16:44:24 +08:00

230 lines
10 KiB
Python

# !/usr/bin/env python
# -*- coding: UTF-8 -*-
import numpy as np
import torch
import os
import folder_paths
from .models.util import load_flow_model,load_ae_wrapper
from .visualcloze_wrapper import VisualClozeModel
from .model_utils import gc_cleanup,phi2narry,pre_x_noise_clip_grid,pre_img_grid,nomarl_upscale_tensor,tensortopil_list_upscale,get_sampler_item,tensortopil_list,load_images_list
from .examples.gradio_tasks import dense_prediction_text,conditional_generation_text
from .examples.gradio_tasks_editing import editing_text
from .examples.gradio_tasks_editing_subject import editing_with_subject_text
from .examples.gradio_tasks_photodoodle import photodoodle_text
from .examples.gradio_tasks_relighting import relighting_text
from .examples.gradio_tasks_restoration import image_restoration_text
from .examples.gradio_tasks_style import style_condition_fusion_text,style_transfer_text
from .examples.gradio_tasks_subject import (subject_driven_text,image_restoration_with_subject_text,
style_transfer_with_subject_text,condition_subject_fusion_text,condition_subject_style_fusion_text)
from .examples.gradio_tasks_tryon import tryon_text
from .examples.gradio_tasks_unseen import unseen_tasks_text
MAX_SEED = np.iinfo(np.int32).max
cur_node_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"
infer_mode_all=dense_prediction_text+conditional_generation_text+editing_text+editing_with_subject_text+photodoodle_text+relighting_text+image_restoration_text+style_condition_fusion_text+style_transfer_text+subject_driven_text+image_restoration_with_subject_text+style_transfer_with_subject_text+condition_subject_fusion_text+condition_subject_style_fusion_text+tryon_text+unseen_tasks_text
class VisualCloze_Aplly:
def __init__(self):
self.counters = {}
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"ckpt": (["none"] + folder_paths.get_filename_list("checkpoints")+folder_paths.get_filename_list("diffusion_models"),),
"lora":(["none"] + folder_paths.get_filename_list("loras"),),
"offload": ("BOOLEAN", {"default": True},), }
}
RETURN_TYPES = ("VisualCloze_PIPE","VisualCloze_INFO")
RETURN_NAMES = ("model","info",)
FUNCTION = "main"
CATEGORY = "VisualCloze"
def main(self, ckpt,lora,offload,):
print("Loading checkpoint...")
if ckpt!="none":
ckpt_path = folder_paths.get_full_path("diffusion_models", ckpt)
else:
raise "ckpt is none"
model_name="flux-dev-fill-lora"
# Load lora model weights
use_lora=True if lora!="none" and "lora" in model_name else False
flow_model = load_flow_model(model_name, ckpt_path,use_lora,offload,device="cpu" if offload else "cuda", lora_rank=256)
resolution=384
if use_lora:
lora_path = folder_paths.get_full_path("loras", lora)
resolution=512 if "512" in lora_path else 384
print(f"Loading lora model from {lora_path}")
ckpt = torch.load(lora_path,weights_only=False,map_location='cpu')
flow_model.load_state_dict(ckpt, strict=False, assign=True)
del ckpt
gc_cleanup()
pipe=VisualClozeModel(flow_model,device)
print("Loading checkpoint is done!")
return (pipe,{"resolution":resolution,"model_name":model_name},)
class VisualCloze_CLIPText:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"vae":("VAE",),
"clip": ("CLIP",),
"info": ("VisualCloze_INFO",),
"query_image":("IMAGE",),
"in_context_img_1":("IMAGE",), #始终开启上下文学习
"infer_tasks": (["none"]+infer_mode_all,),
"seed": ("INT", {"default": 0, "min": 0, "max": MAX_SEED,}),
# "width": ("INT", {"default": 768, "min": 256, "max": 2048, "step": 16, "display": "number"}),
# "height": ("INT", {"default": 768, "min": 256, "max": 2048, "step": 16, "display": "number"}),
"content_prompt":("STRING", {"multiline": True,"default": ""}),
},
"optional": {
"in_context_img_2":("IMAGE",),
"in_context_img_3":("IMAGE",),
"in_context_img_4":("IMAGE",),
}
}
RETURN_TYPES = ("VisualCloze_EMB",)
RETURN_NAMES = ("emb_dict",)
FUNCTION = "encode"
CATEGORY = "VisualCloze"
def encode(self, vae,clip,info, query_image,in_context_img_1,infer_tasks,seed, content_prompt,**kwargs):
# use normal vae
vae=load_ae_wrapper(info.get("model_name"),vae.get_sd())
vae.requires_grad_(False)
assert infer_tasks!="none","please choice a task"
# pre imge and get grid
# in_pil_img_1=tensortopil_list_upscale(kwargs.get("in_context_img_1"),width,height) if isinstance(kwargs.get("in_context_img_1"),torch.Tensor) else None
# in_pil_img_2=tensortopil_list_upscale(kwargs.get("in_context_img_2"),width,height) if isinstance(kwargs.get("in_context_img_2"),torch.Tensor) else None
# in_pil_img_3=tensortopil_list_upscale(kwargs.get("in_context_img_3"),width,height) if isinstance(kwargs.get("in_context_img_3"),torch.Tensor) else None
in_pil_img_2=tensortopil_list(kwargs.get("in_context_img_2")) if isinstance(kwargs.get("in_context_img_2"),torch.Tensor) else None
in_pil_img_3=tensortopil_list(kwargs.get("in_context_img_3")) if isinstance(kwargs.get("in_context_img_3"),torch.Tensor) else None
in_pil_img_4=tensortopil_list(kwargs.get("in_context_img_4")) if isinstance(kwargs.get("in_context_img_4"),torch.Tensor) else None
grid_imag_list,grid_w,grid_h=pre_img_grid(query_image, in_context_img_1, in_pil_img_2, in_pil_img_3, in_pil_img_4)
print(grid_imag_list,grid_w,grid_h)
outputs=get_sampler_item(infer_tasks,grid_w,grid_h,content_prompt)
generator = torch.Generator(device=device).manual_seed(int(seed))
# clip
inp,img_cond,sliced_subimage,mask_position,upsampling_size = pre_x_noise_clip_grid(clip,vae,generator,grid_imag_list,
[outputs[1]+" "+ outputs[2]+" "+outputs[3]], grid_h,grid_w,info.get("resolution"),device,dtype=torch.bfloat16)
emb_dict={"sliced_subimage":sliced_subimage, "mask_position":mask_position,"inp":inp,"img_cond":img_cond,"clip":clip,"grid_h":grid_h,"grid_w":grid_w,
"content_prompt":outputs[3],"generator":generator,"vae":vae,"upsampling_size":upsampling_size,"upsampling_noise":outputs[4],"steps":outputs[5]}
return (emb_dict,)
class VisualCloze_KSampler:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("VisualCloze_PIPE",),
"emb_dict": ("VisualCloze_EMB",),
"steps": ("INT", {"default": 30, "min": 1, "max": 10000,}),
"upsampling_steps": ("INT", {"default": 10, "min": 1, "max": 10000,}),
"upsampling_noise": ("FLOAT", {"default": 0.4, "min": 0.0, "max": 1.0, "step": 0.1,}),
"cfg": ("FLOAT", {"default": 30.0, "min": 0.0, "max": 100.0, "step": 0.1,}),
},
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
FUNCTION = "sample"
CATEGORY = "VisualCloze"
def sample(self, model,emb_dict, steps,upsampling_steps,upsampling_noise, cfg,):
result = model.process_images(
emb_dict.get("vae"),
emb_dict.get("clip"),
cfg,
steps=steps if emb_dict.get("steps") is None else emb_dict.get("steps"),
upsampling_steps=upsampling_steps,
upsampling_noise=upsampling_noise if emb_dict.get("upsampling_noise") is None else emb_dict.get("upsampling_noise"),
inp=emb_dict.get("inp"),
img_cond=emb_dict.get("img_cond"),
sliced_subimage=emb_dict.get("sliced_subimage"),
mask_position=emb_dict.get("mask_position"),
upsampling_size=emb_dict.get("upsampling_size"),
content_prompt=emb_dict.get("content_prompt"),
generator=emb_dict.get("generator"),
grid_h=emb_dict.get("grid_h"),
grid_w=emb_dict.get("grid_w"),
)[-1]
return (phi2narry(result),)
class Img_Quadruple:
@classmethod
def INPUT_TYPES(s):
return {"required": {"image_1": ("IMAGE",), # [B,H,W,C], C=3
},
"optional": {"image_2": ("IMAGE",),
"image_3": ("IMAGE",),
"image_4": ("IMAGE",)}
}
RETURN_TYPES = ("IMAGE",)
ETURN_NAMES = ("image",)
FUNCTION = "main"
CATEGORY = "VisualCloze"
def main(self, image_1, **kwargs):
B,height,width, _ = image_1.size()
assert B == 1 , "only support batch size 1"
image_2 = nomarl_upscale_tensor(kwargs.get("image_2"), width, height) if isinstance(kwargs.get("image_2"), torch.Tensor) else None
image_3 = nomarl_upscale_tensor(kwargs.get("image_3"), width, height) if isinstance(kwargs.get("image_3"), torch.Tensor) else None
image_4 = nomarl_upscale_tensor(kwargs.get("image_4"), width, height) if isinstance(kwargs.get("image_4"), torch.Tensor) else None
img_list = [image_1]
for img in [image_2, image_3, image_4]:
if img is not None:
C,_,_, _ = img.size()
assert C == 1 , "only support batch size 1"
img_list.append(img)
images = torch.cat(tuple(img_list), dim=0)
return (images,)
NODE_CLASS_MAPPINGS = {
"VisualCloze_Aplly": VisualCloze_Aplly,
"VisualCloze_KSampler": VisualCloze_KSampler,
"VisualCloze_CLIPText": VisualCloze_CLIPText,
"Img_Quadruple": Img_Quadruple,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"MSdiffusion_Aplly": "MSdiffusion_Aplly",
"VisualCloze_KSampler": "VisualCloze_KSampler",
"VisualCloze_CLIPText": "VisualCloze_CLIPText",
"Img_Quadruple": "Img_Quadruple",
}