moralec/MagicQuill
0
1import torch.nn.functional as F2import torch3import numpy as np4from PIL import Image5import os6import sys7sys.path.append(os.path.dirname(os.path.abspath(__file__)))8 9from .brushnet_nodes import BrushNetLoader, BrushNet, BlendInpaint, get_files_with_extension10from .comfyui_utils import CheckpointLoaderSimple, ControlNetLoader, ControlNetApplyAdvanced, CLIPTextEncode, KSampler, VAEDecode, GrowMask, PIDINET_Preprocessor, LineArt_Preprocessor, Color_Preprocessor11 12class ScribbleColorEditModel():13 def __init__(self):14 self.checkpoint_loader = CheckpointLoaderSimple()15 self.clip_text_encoder = CLIPTextEncode()16 self.mask_processor = GrowMask()17 self.controlnet_loader = ControlNetLoader()18 self.scribble_processor = PIDINET_Preprocessor()19 self.lineart_processor = LineArt_Preprocessor()20 self.color_processor = Color_Preprocessor()21 self.brushnet_loader = BrushNetLoader()22 self.brushnet_node = BrushNet()23 self.controlnet_apply = ControlNetApplyAdvanced()24 self.ksampler = KSampler()25 self.vae_decoder = VAEDecode()26 self.blender = BlendInpaint()27 self.ckpt_name = os.path.normpath("SD1.5/realisticVisionV60B1_v51VAE.safetensors")28 with torch.no_grad():29 self.model, self.clip, self.vae = self.checkpoint_loader.load_checkpoint(self.ckpt_name)30 self.load_models('SD1.5', 'float16')31 32 def load_models(self, base_model_version="SD1.5", dtype='float16'):33 if base_model_version == "SD1.5":34 edge_controlnet_name = "control_v11p_sd15_scribble.safetensors"35 color_controlnet_name = "color_finetune.safetensors"36 brushnet_name = os.path.normpath("brushnet/random_mask_brushnet_ckpt/diffusion_pytorch_model.safetensors")37 else:38 raise ValueError("Invalid base_model_version, not supported yet!!!: {}".format(base_model_version))39 self.edge_controlnet = self.controlnet_loader.load_controlnet(edge_controlnet_name)[0]40 self.color_controlnet = self.controlnet_loader.load_controlnet(color_controlnet_name)[0]41 self.brushnet_loader.inpaint_files = get_files_with_extension('inpaint')42 print("self.brushnet_loader.inpaint_files: ", get_files_with_extension('inpaint'))43 self.brushnet = self.brushnet_loader.brushnet_loading(brushnet_name, dtype)[0]44 45 def process(self, ckpt_name, image, colored_image, positive_prompt, negative_prompt, mask, add_mask, remove_mask, grow_size, stroke_as_edge, fine_edge, edge_strength, color_strength, inpaint_strength, seed, steps, cfg, sampler_name, scheduler, base_model_version='SD1.5', dtype='float16', palette_resolution=2048):46 if ckpt_name != self.ckpt_name:47 self.ckpt_name = ckpt_name48 with torch.no_grad():49 self.model, self.clip, self.vae = self.checkpoint_loader.load_checkpoint(ckpt_name)50 if not hasattr(self, 'edge_controlnet') or not hasattr(self, 'color_controlnet') or not hasattr(self, 'brushnet'):51 self.load_models(base_model_version, dtype)52 positive = self.clip_text_encoder.encode(self.clip, positive_prompt)[0]53 negative = self.clip_text_encoder.encode(self.clip, negative_prompt)[0] 54 # Grow Mask for Color Editing55 mask = self.mask_processor.expand_mask(mask, expand=grow_size, tapered_corners=True)[0]56 # Realistic Lineart57 image_copy = image.clone()58 if stroke_as_edge == "disable":59 bool_add_mask = add_mask > 0.560 mean_brightness = image_copy[bool_add_mask].mean()61 if mean_brightness > 0.8:62 image_copy[bool_add_mask] = 0.063 else:64 image_copy[bool_add_mask] = 1.065 66 67 if not torch.equal(image, colored_image):68 print("Apply color controlnet")69 color_output = self.color_processor.execute(colored_image, resolution=palette_resolution)[0]70 lineart_output = self.lineart_processor.execute(image, resolution=512, coarse=False)[0]71 positive, negative = self.controlnet_apply.apply_controlnet(positive, negative, self.color_controlnet, color_output, color_strength, 0.0, 1.0)72 positive, negative = self.controlnet_apply.apply_controlnet(positive, negative, self.edge_controlnet, lineart_output, 0.8, 0.0, 1.0)73 else:74 print("Apply edge controlnet")75 # Resize masks to match the dimensions of lineart_output76 color_output = self.color_processor.execute(image, resolution=palette_resolution)[0]77 if fine_edge == "enable":78 lineart_output = self.lineart_processor.execute(image, resolution=512, coarse=False)[0]79 else:80 lineart_output = self.scribble_processor.execute(image, resolution=512)[0]81 add_mask_resized = F.interpolate(add_mask.unsqueeze(0).unsqueeze(0).float(), size=(1, lineart_output.shape[1], lineart_output.shape[2]), mode='nearest').squeeze(0).squeeze(0)82 remove_mask_resized = F.interpolate(remove_mask.unsqueeze(0).unsqueeze(0).float(), size=(1, lineart_output.shape[1], lineart_output.shape[2]), mode='nearest').squeeze(0).squeeze(0)83 84 bool_add_mask_resized = (add_mask_resized > 0.5)85 bool_remove_mask_resized = (remove_mask_resized > 0.5)86 87 if stroke_as_edge == "enable":88 lineart_output[bool_remove_mask_resized] = 0.089 lineart_output[bool_add_mask_resized] = 1.090 else:91 lineart_output[bool_remove_mask_resized & ~bool_add_mask_resized] = 0.092 positive, negative = self.controlnet_apply.apply_controlnet(positive, negative, self.edge_controlnet, lineart_output, edge_strength, 0.0, 1.0)93 94 # BrushNet95 model, positive, negative, latent = self.brushnet_node.model_update(96 model=self.model,97 vae=self.vae,98 image=image,99 mask=mask,100 brushnet=self.brushnet,101 positive=positive,102 negative=negative,103 scale=inpaint_strength,104 start_at=0,105 end_at=10000106 )107 108 # KSampler Node109 latent_samples = self.ksampler.sample(110 model=model, 111 seed=seed, 112 steps=steps, 113 cfg=cfg, 114 sampler_name=sampler_name, 115 scheduler=scheduler, 116 positive=positive, 117 negative=negative, 118 latent_image=latent,119 )[0]120 121 final_image = self.vae_decoder.decode(self.vae, latent_samples)[0]122 final_image = self.blender.blend_inpaint(final_image, image, mask, kernel=10, sigma=10.0)[0]123 124 # Return the final image125 return (latent_samples, final_image, lineart_output, color_output)126 