CoolFace
Apppublic

moralec/MagicQuill

sourceHugging Facecc-by-nc-4.0updated 2y agoView on Hugging Face
0likes
scribble_color_edit.py126 linesDownload Raw Back to MagicQuill
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