Add files via upload
This commit is contained in:
@@ -0,0 +1,176 @@
|
||||
import torch
|
||||
import gc
|
||||
import os
|
||||
import numpy as np
|
||||
import json
|
||||
import comfy.model_management as mm
|
||||
from PIL import Image
|
||||
|
||||
from folder_paths import map_legacy, folder_names_and_paths
|
||||
from .controlnet_union import ControlNetModel_Union
|
||||
from .pipeline_fill_sd_xl import encode_prompt, StableDiffusionXLFillPipeline
|
||||
from diffusers import AutoencoderKL, TCDScheduler
|
||||
from diffusers.models.model_loading_utils import load_state_dict
|
||||
from transformers import CLIPTextModel, CLIPTextModelWithProjection, CLIPTokenizer
|
||||
|
||||
|
||||
def get_first_folder_list(folder_name: str) -> tuple[list[str], dict[str, float], float]:
|
||||
folder_name = map_legacy(folder_name)
|
||||
global folder_names_and_paths
|
||||
folders = folder_names_and_paths[folder_name]
|
||||
if folder_name == "unet":
|
||||
root_folder = folders[0][0]
|
||||
elif folder_name == "diffusion_models":
|
||||
root_folder = folders[0][1]
|
||||
visible_folders = [name for name in os.listdir(root_folder) if os.path.isdir(os.path.join(root_folder, name))]
|
||||
return visible_folders
|
||||
|
||||
|
||||
# Tensor to PIL (grabbed from WAS Suite)
|
||||
def tensor2pil(image: torch.Tensor) -> Image.Image:
|
||||
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
|
||||
|
||||
# Convert PIL to Tensor (grabbed from WAS Suite)
|
||||
def pil2tensor(image: Image.Image) -> torch.Tensor:
|
||||
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
|
||||
|
||||
def get_device_by_name(device):
|
||||
if device == 'auto':
|
||||
device = mm.get_torch_device()
|
||||
return device
|
||||
|
||||
|
||||
def get_dtype_by_name(dtype):
|
||||
if dtype == 'auto':
|
||||
if mm.should_use_fp16():
|
||||
dtype = torch.float16
|
||||
elif mm.should_use_bf16():
|
||||
dtype = torch.bfloat16
|
||||
else:
|
||||
dtype = torch.float32
|
||||
elif dtype== "fp16":
|
||||
dtype = torch.float16
|
||||
elif dtype == "bf16":
|
||||
dtype = torch.bfloat16
|
||||
elif dtype == "fp32":
|
||||
dtype = torch.float32
|
||||
elif dtype == "fp8_e4m3fn":
|
||||
dtype = torch.float8_e4m3fn
|
||||
elif dtype == "fp8_e4m3fnuz":
|
||||
dtype = torch.float8_e4m3fnuz
|
||||
elif dtype == "fp8_e5m2":
|
||||
dtype = torch.float8_e5m2
|
||||
elif dtype == "fp8_e5m2fnuz":
|
||||
dtype = torch.float8_e5m2fnuz
|
||||
|
||||
return dtype
|
||||
|
||||
|
||||
def loadDiffModels1(model_path, dtype, device):
|
||||
tokenizer = CLIPTokenizer.from_pretrained(model_path, subfolder="tokenizer", use_fast=False)
|
||||
tokenizer_2 = CLIPTokenizer.from_pretrained(model_path, subfolder="tokenizer_2", use_fast=False)
|
||||
text_encoder = CLIPTextModel.from_pretrained(model_path, subfolder="text_encoder", torch_dtype=dtype).requires_grad_(False).to(device)
|
||||
text_encoder_2 = CLIPTextModelWithProjection.from_pretrained(model_path, subfolder="text_encoder_2", torch_dtype=dtype).requires_grad_(False).to(device)
|
||||
|
||||
return tokenizer, tokenizer_2, text_encoder, text_encoder_2
|
||||
|
||||
|
||||
def encodeDiffOutpaintPrompt(model_path, dtype, final_prompt, device):
|
||||
tokenizer, tokenizer_2, text_encoder, text_encoder_2 = loadDiffModels1(model_path, dtype, device)
|
||||
|
||||
(prompt_embeds,
|
||||
negative_prompt_embeds,
|
||||
pooled_prompt_embeds,
|
||||
negative_pooled_prompt_embeds,
|
||||
) = encode_prompt(final_prompt, tokenizer, tokenizer_2, text_encoder, text_encoder_2, device, True)
|
||||
|
||||
del tokenizer, tokenizer_2, text_encoder, text_encoder_2
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
|
||||
return prompt_embeds, negative_prompt_embeds, pooled_prompt_embeds, negative_pooled_prompt_embeds
|
||||
|
||||
|
||||
def loadControlnetModel(device, dtype, controlnet_path):
|
||||
config_file = f"{controlnet_path}/config_promax.json"
|
||||
config = ControlNetModel_Union.load_config(config_file)
|
||||
controlnet_model = ControlNetModel_Union.from_config(config)
|
||||
|
||||
model_file = f"{controlnet_path}/diffusion_pytorch_model_promax.safetensors"
|
||||
state_dict = load_state_dict(model_file)
|
||||
|
||||
model, _, _, _, _ = ControlNetModel_Union._load_pretrained_model(
|
||||
controlnet_model, state_dict, model_file, f"{controlnet_path}"
|
||||
)
|
||||
controlnet_model.to(device, dtype)
|
||||
|
||||
del model, state_dict, model_file
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
|
||||
return controlnet_model
|
||||
|
||||
|
||||
def loadVaeModel(vae_path, device, dtype, enable_vae_slicing, enable_vae_tiling):
|
||||
vae = AutoencoderKL.from_pretrained(f"{vae_path}").to(device, dtype)
|
||||
if enable_vae_slicing:
|
||||
vae.enable_slicing()
|
||||
else:
|
||||
vae.disable_slicing()
|
||||
|
||||
if enable_vae_tiling:
|
||||
vae.enable_tiling()
|
||||
else:
|
||||
vae.disable_tiling()
|
||||
return vae
|
||||
|
||||
|
||||
def diffuserOutpaintSamples(model_path, controlnet_model, diffuser_outpaint_cnet_image, dtype, controlnet_path,
|
||||
prompt_embeds, negative_prompt_embeds, pooled_prompt_embeds, negative_pooled_prompt_embeds,
|
||||
device, steps, controlnet_strength, guidance_scale,
|
||||
keep_model_device):
|
||||
|
||||
controlnet_model = loadControlnetModel(device, dtype, controlnet_path)
|
||||
|
||||
with open(f"{model_path}/scheduler/scheduler_config.json", "r") as f:
|
||||
scheduler_config = json.load(f)
|
||||
scheduler = TCDScheduler.from_config(scheduler_config)
|
||||
|
||||
pipe = StableDiffusionXLFillPipeline.from_pretrained(
|
||||
model_path,
|
||||
torch_dtype=dtype,
|
||||
variant="fp16",
|
||||
scheduler=scheduler,
|
||||
)
|
||||
if not keep_model_device:
|
||||
pipe.to(device)
|
||||
|
||||
cnet_image = diffuser_outpaint_cnet_image
|
||||
cnet_image=tensor2pil(cnet_image)
|
||||
cnet_image=cnet_image.convert('RGB')
|
||||
|
||||
rgb_latents = list(pipe(
|
||||
prompt_embeds=prompt_embeds,
|
||||
negative_prompt_embeds=negative_prompt_embeds,
|
||||
pooled_prompt_embeds=pooled_prompt_embeds,
|
||||
negative_pooled_prompt_embeds=negative_pooled_prompt_embeds,
|
||||
image=cnet_image,
|
||||
num_inference_steps=steps,
|
||||
controlnet_model=controlnet_model,
|
||||
controlnet_conditioning_scale=controlnet_strength,
|
||||
guidance_scale=guidance_scale,
|
||||
device=device,
|
||||
keep_model_device=keep_model_device,
|
||||
))
|
||||
|
||||
last_rgb_latent = rgb_latents[-1] # Access the last image
|
||||
|
||||
del pipe, controlnet_model, scheduler, prompt_embeds, negative_prompt_embeds, pooled_prompt_embeds, negative_pooled_prompt_embeds
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
|
||||
return last_rgb_latent
|
||||
Reference in New Issue
Block a user