Team Ai
Apppublic

Milcho/ControlNet-Guidance

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
app.py131 linesDownload Raw Back to root
1from share import *2import config3import os4 5import cv26import einops7import gradio as gr8import numpy as np9import torch10import random11 12from huggingface_hub import hf_hub_download13from pytorch_lightning import seed_everything14from annotator.util import resize_image, HWC315from annotator.hed import HEDdetector16from cldm.model import create_model, load_state_dict17from cldm.ddim_hacked import DDIMSampler18 19from PIL import Image20import clip21 22 23device = "cuda" if torch.cuda.is_available() else "cpu"24 25if torch.cuda.is_available():26    print("\nCUDA (GPU support) is available in PyTorch!")27    print(f"Number of GPUs available: {torch.cuda.device_count()}")28    print(f"GPU Name: {torch.cuda.get_device_name(0)}\n")29else:30    print("\nCUDA (GPU support) is not available in PyTorch.\n")31 32# Function to download the model from Hugging Face Hub33def download_model_if_not_exists(repo_id, filename, cache_dir="./models"):34    model_path = os.path.join(cache_dir, filename)35    if not os.path.exists(model_path):36        print(f"Model file not found. Downloading to {model_path}...")37        model_path = hf_hub_download(repo_id=repo_id, filename=filename, cache_dir=cache_dir)38    else:39        print(f"Model file found at {model_path}. Using the existing file.")40    return model_path41 42# Specify the model path and YAML config path43yaml_path = "./models/cldm_v15.yaml"44model_filename = "control_sd15_hed.pth"45 46# Download the necessary model file if it doesn't exist47model_path = download_model_if_not_exists("lllyasviel/ControlNet", model_filename)48 49# Load the CLDM model50model = create_model(yaml_path).to(device)51model.load_state_dict(load_state_dict(model_path, location=device))52model = model.cuda()53ddim_sampler = DDIMSampler(model)54 55def assess_sketch_quality(image, prompt, clip_model, preprocess):56    image = Image.fromarray(image)57    text_tokens = clip.tokenize([prompt]).to(device)58    image_preprocessed = preprocess(image).unsqueeze(0).to(device)59 60    with torch.no_grad():61        image_features = clip_model.encode_image(image_preprocessed)62        text_features = clip_model.encode_text(text_tokens)63 64    similarity = torch.nn.functional.softmax((image_features @ text_features.T).squeeze(0), dim=0)65    quality_score = similarity.item()66 67    quality_score = quality_score * 0.868 69    return quality_score70 71def process(input_image, prompt, a_prompt, n_prompt, num_samples, image_resolution, detect_resolution, ddim_steps, guess_mode, strength, scale, seed, eta):72    with torch.no_grad():73        input_image = HWC3(input_image)74        quality_score = assess_sketch_quality(input_image, prompt, clip_model, preprocess)75        strength = quality_score76 77        detected_map = apply_hed(resize_image(input_image, detect_resolution))78        detected_map = HWC3(detected_map)79        img = resize_image(input_image, image_resolution)80        H, W, C = img.shape81        detected_map = cv2.resize(detected_map, (W, H), interpolation=cv2.INTER_LINEAR)82 83        control = torch.from_numpy(detected_map.copy()).float().to(device) / 255.084        control = torch.stack([control for _ in range(num_samples)], dim=0)85        control = einops.rearrange(control, 'b h w c -> b c h w').clone()86 87        if seed == -1:88            seed = random.randint(0, 65535)89        seed_everything(seed)90 91        cond = {"c_concat": [control], "c_crossattn": [model.get_learned_conditioning([prompt + ', ' + a_prompt] * num_samples)]}92        un_cond = {"c_concat": None if guess_mode else [control], "c_crossattn": [model.get_learned_conditioning([n_prompt] * num_samples)]}93        shape = (4, H // 8, W // 8)94 95        model.control_scales = [strength * (0.825 ** float(12 - i)) for i in range(13)] if guess_mode else ([strength] * 13)96        samples, intermediates = ddim_sampler.sample(ddim_steps, num_samples, shape, cond, verbose=False, eta=eta, unconditional_guidance_scale=scale, unconditional_conditioning=un_cond)97 98        x_samples = model.decode_first_stage(samples)99        x_samples = (einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 + 127.5).cpu().numpy().clip(0, 255).astype(np.uint8)100 101        results = [x_samples[i] for i in range(num_samples)]102    return [detected_map] + results103 104block = gr.Blocks().queue()105with block:106    with gr.Row():107        gr.Markdown("## ControlNet Integrated with CLIP")108    with gr.Row():109        with gr.Column():110            input_image = gr.Image(source='upload', type="numpy")111            prompt = gr.Textbox(label="Describe the Image")112            run_button = gr.Button(label="Run")113            with gr.Accordion("Advanced options", open=False):114                num_samples = gr.Slider(label="Images", minimum=1, maximum=12, value=1, step=1)115                image_resolution = gr.Slider(label="Image Resolution", minimum=256, maximum=768, value=512, step=64)116                strength = gr.Slider(label="Control Strength", minimum=0.0, maximum=2.0, value=1.0, step=0.01)117                guess_mode = gr.Checkbox(label='Guess Mode', value=False)118                detect_resolution = gr.Slider(label="HED Resolution", minimum=128, maximum=1024, value=512, step=1)119                ddim_steps = gr.Slider(label="Steps", minimum=1, maximum=100, value=20, step=1)120                scale = gr.Slider(label="Guidance Scale", minimum=0.1, maximum=30.0, value=9.0, step=0.1)121                seed = gr.Slider(label="Seed", minimum=-1, maximum=2147483647, step=1, randomize=True)122                eta = gr.Number(label="eta (DDIM)", value=0.0)123                a_prompt = gr.Textbox(label="Added Prompt", value='best quality, extremely detailed')124                n_prompt = gr.Textbox(label="Negative Prompt", value='longbody, lowres, bad anatomy, bad hands, missing fingers, extra digit, fewer digits, cropped, worst quality, low quality')125        with gr.Column():126            result_gallery = gr.Gallery(label='Output', show_label=False, elem_id="gallery").style(grid=2, height='auto')127    ips = [input_image, prompt, a_prompt, n_prompt, num_samples, image_resolution, detect_resolution, ddim_steps, guess_mode, strength, scale, seed, eta]128    run_button.click(fn=process, inputs=ips, outputs=[result_gallery])129 130block.launch()131