Team Ai
Apppublic

dezzman/diffusion_models

sourceHugging Faceupdated 2y agoView on Hugging Face
2likes
app.py330 linesDownload Raw Back to root
1import gradio as gr2import numpy as np3import torch4from diffusers.utils import load_image5from diffusers import StableDiffusionControlNetPipeline, ControlNetModel6from peft import PeftModel, LoraConfig7from controlnet_aux import HEDdetector8from PIL import Image9import cv2 as cv10import os11from functools import lru_cache12from contextlib import contextmanager13 14MAX_SEED = np.iinfo(np.int32).max15MAX_IMAGE_SIZE = 102416IP_ADAPTER = 'h94/IP-Adapter'17IP_ADAPTER_WEIGHT_NAME = "ip-adapter-plus_sd15.bin"18 19device = torch.device("cuda" if torch.cuda.is_available() else "cpu")20model_id_default = "stable-diffusion-v1-5/stable-diffusion-v1-5"21torch_dtype = torch.float16 if torch.cuda.is_available() else torch.float3222 23class PipelineManager:24    def __init__(self):25        self.pipe = None26        self.current_model = None27        self.controlnet_cache = {}28        self.hed = None29        30    @lru_cache(maxsize=2)31    def get_controlnet(self, model_name: str) -> ControlNetModel:32        if model_name not in self.controlnet_cache:33            self.controlnet_cache[model_name] = ControlNetModel.from_pretrained(34                model_name, 35                cache_dir="./models_cache",36                torch_dtype=torch_dtype37            ).to(device)38        return self.controlnet_cache[model_name]39    40    def get_hed_detector(self):41        if self.hed is None:42            self.hed = HEDdetector.from_pretrained('lllyasviel/Annotators')43        return self.hed44 45    def initialize_pipeline(self, model_id, controlnet_model):46        controlnet = self.get_controlnet(controlnet_model)47        if not self.pipe or model_id != self.current_model:48            self.pipe = self.create_pipeline(model_id, controlnet)49            self.current_model = model_id50        return self.pipe51 52    def create_pipeline(self, model_id, controlnet):53        pipe = StableDiffusionControlNetPipeline.from_pretrained(54            model_id,55            torch_dtype=torch_dtype,56            controlnet=controlnet,57            cache_dir="./models_cache"58        ).to(device)59        60        if os.path.exists('./lora_logos'):61            pipe = self.load_lora_adapters(pipe)62            63        return pipe64 65    def load_lora_adapters(self, pipe):66        unet_dir = os.path.join('./lora_logos', "unet")67        text_encoder_dir = os.path.join('./lora_logos', "text_encoder")68        69        pipe.unet = PeftModel.from_pretrained(pipe.unet, unet_dir, adapter_name="default")70        if os.path.exists(text_encoder_dir):71            pipe.text_encoder = PeftModel.from_pretrained(pipe.text_encoder, text_encoder_dir)72            73        return pipe.to(device)74 75@contextmanager76def torch_inference_mode():77    with torch.inference_mode(), torch.autocast(device.type):78        yield79 80def process_embeddings(prompt, negative_prompt, tokenizer, text_encoder):81    def process_text(text):82        tokens = tokenizer(text, return_tensors="pt", truncation=False).input_ids83        chunks = [tokens[:, i:i+77].to(device) for i in range(0, tokens.size(1), 77)]84        return torch.cat([text_encoder(chunk)[0] for chunk in chunks], dim=1)85    86    prompt_emb = process_text(prompt)87    negative_emb = process_text(negative_prompt)88    max_len = max(prompt_emb.size(1), negative_emb.size(1))89    90    return (91        torch.nn.functional.pad(prompt_emb, (0, 0, 0, max_len - prompt_emb.size(1))),92        torch.nn.functional.pad(negative_emb, (0, 0, 0, max_len - negative_emb.size(1)))93    )94 95def process_control_image(image_path: str, processor: str, hed_detector) -> Image:96    image = load_image(image_path).convert('RGB')97    98    if processor == 'edge_detection':99        edges = cv.Canny(np.array(image), 80, 160)100        return Image.fromarray(np.repeat(edges[:, :, None], 3, axis=2))101    102    if processor == 'scribble':103        scribble = hed_detector(image)104        processed = cv.medianBlur(np.array(scribble), 3)105        return Image.fromarray(cv.convertScaleAbs(processed, alpha=1.5))106 107pipeline_mgr = PipelineManager()108controlnet_models = {109    "edge_detection": "lllyasviel/sd-controlnet-canny",110    "scribble": "lllyasviel/sd-controlnet-scribble"111}112 113def infer(114    prompt, 115    negative_prompt, 116    width=512, 117    height=512, 118    num_inference_steps=20, 119    model_id='stable-diffusion-v1-5/stable-diffusion-v1-5', 120    seed=42, 121    guidance_scale=7.0, 122    lora_scale=0.5,123    cn_enable=False,124    cn_strength=0.0,125    cn_mode='edge_detection',126    cn_image=None,127    ip_enable=False,128    ip_scale=0.5,129    ip_image=None,130    progress=gr.Progress(track_tqdm=True)131    ):132 133    generator = torch.Generator(device).manual_seed(seed)134    135    with torch_inference_mode():136        pipe = pipeline_mgr.initialize_pipeline(137            model_id, 138            controlnet_models.get(cn_mode, controlnet_models['edge_detection'])139        )140        141        if cn_enable and not cn_image:142            raise gr.Error("ControlNet enabled but no image provided!")143 144        if ip_enable and not ip_image:145            raise gr.Error("IP-Adapter enabled but no image provided!")146        147        prompt_emb, negative_emb = process_embeddings(148            prompt, 149            negative_prompt, 150            pipe.tokenizer, 151            pipe.text_encoder152        )153        154        params = {155            'prompt_embeds': prompt_emb,156            'negative_prompt_embeds': negative_emb,157            'guidance_scale': guidance_scale,158            'num_inference_steps': num_inference_steps,159            'width': width,160            'height': height,161            'generator': generator,162            'cross_attention_kwargs': {"scale": lora_scale},163        }164        165        if cn_enable:166            params['image'] = process_control_image(167                cn_image,168                cn_mode,169                pipeline_mgr.get_hed_detector()170            )171            params['controlnet_conditioning_scale'] = float(cn_strength)172        else:173            params['image'] = torch.zeros((1, 3, 512, 512)).to(device)  # заглушка, чтобы pipeline не падал174            params['controlnet_conditioning_scale'] = 0.0175            176        if ip_enable:177            pipe.load_ip_adapter(IP_ADAPTER, subfolder="models", weight_name=IP_ADAPTER_WEIGHT_NAME)178            params['ip_adapter_image'] = load_image(ip_image).convert('RGB')179            pipe.set_ip_adapter_scale(ip_scale)180            181        pipe.fuse_lora(lora_scale=lora_scale)182        183        return pipe(**params).images[0]184 185css = """186#col-container {187    margin: 0 auto;188    max-width: 640px;189}190"""191 192with gr.Blocks(css=css) as demo:193    with gr.Column(elem_id="col-container"):194        gr.Markdown("# ⚽️ Football Logo Generator")195        196        with gr.Row():197            model_id = gr.Textbox(198                label="Model ID",199                max_lines=1,200                placeholder="Enter model id like 'stable-diffusion-v1-5/stable-diffusion-v1-5'",201                value=model_id_default202            )203 204        prompt = gr.Textbox(205            label="Prompt",206            max_lines=1,207            placeholder="Enter your prompt",208        )209 210        negative_prompt = gr.Textbox(211            label="Negative prompt",212            max_lines=1,213            placeholder="Enter a negative prompt",214        )215 216        with gr.Row():217            seed = gr.Number(218                label="Seed",219                minimum=0,220                maximum=MAX_SEED,221                step=1,222                value=42,223            )224 225        with gr.Row():226            guidance_scale = gr.Slider(227                label="Guidance scale",228                minimum=0.0,229                maximum=10.0,230                step=0.1,231                value=7.0,232            )233 234        with gr.Row():235            lora_scale = gr.Slider(236                label="LoRA scale",237                minimum=0.0,238                maximum=1.0,239                step=0.1,240                value=0.5,241            )242 243        with gr.Row():244            num_inference_steps = gr.Slider(245                label="Number of inference steps",246                minimum=1,247                maximum=50,248                step=1,249                value=20,250            )251 252        # Секция Control Net253        cn_enable = gr.Checkbox(label="Enable ControlNet")    254        with gr.Column(visible=False) as cn_options:255            with gr.Row():256                cn_strength = gr.Slider(0, 2, value=0.8, step=0.1, label="Control strength", interactive=True)257                cn_mode = gr.Dropdown(258                    choices=["edge_detection", "scribble"],259                    value="edge_detection",260                    label="Work regime",261                    interactive=True,262                )263            cn_image = gr.Image(type="filepath", label="Control image")264 265        cn_enable.change(266            lambda x: gr.update(visible=x),267            inputs=cn_enable,268            outputs=cn_options269        )270        271        # Секция IP-Adapter272        ip_enable = gr.Checkbox(label="Enable IP-Adapter")273        with gr.Column(visible=False) as ip_options:274            ip_scale = gr.Slider(0, 1, value=0.5, step=0.1, label="IP-adapter scale", interactive=True)275            ip_image = gr.Image(type="filepath", label="IP-adapter image", interactive=True)276 277        ip_enable.change(278            lambda x: gr.update(visible=x),279            inputs=ip_enable,280            outputs=ip_options281        )282 283        with gr.Accordion("Optional Settings", open=False):284            with gr.Row():285                width = gr.Slider(286                    label="Width",287                    minimum=256,288                    maximum=MAX_IMAGE_SIZE,289                    step=32,290                    value=512,291                )292            293            with gr.Row():294                height = gr.Slider(295                    label="Height",296                    minimum=256,297                    maximum=MAX_IMAGE_SIZE,298                    step=32,299                    value=512,300                )301 302        run_button = gr.Button("Run", scale=1, variant="primary")303        result = gr.Image(label="Result", show_label=False)304    305    gr.on(306        triggers=[run_button.click, prompt.submit],307        fn=infer,308        inputs=[309            prompt,310            negative_prompt,311            width,312            height,313            num_inference_steps,314            model_id,315            seed,316            guidance_scale,317            lora_scale,318            cn_enable,319            cn_strength,320            cn_mode,321            cn_image,322            ip_enable,323            ip_scale,324            ip_image325        ],326        outputs=[result],327    )328 329if __name__ == "__main__":330    demo.launch()