Team Ai
Apppublic

SemaSci/DiffModels

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
app.py327 linesDownload Raw Back to root
1import gradio as gr2import numpy as np3import random4import os5import torch6from diffusers import StableDiffusionPipeline, ControlNetModel, StableDiffusionControlNetPipeline7from diffusers.utils import load_image8from peft import PeftModel, LoraConfig9 10from rembg import remove 11 12 13device = "cuda" if torch.cuda.is_available() else "cpu"14model_id_default = "stable-diffusion-v1-5/stable-diffusion-v1-5"15 16if torch.cuda.is_available():17    torch_dtype = torch.float1618else:19    torch_dtype = torch.float3220 21MAX_SEED = np.iinfo(np.int32).max22MAX_IMAGE_SIZE = 102423 24 25# @spaces.GPU #[uncomment to use ZeroGPU]26def infer(27    prompt,28    negative_prompt,29    width=512,30    height=512,31    model_id=model_id_default,32    seed=42,33    guidance_scale=7.0,34    lora_scale=1.0,35    num_inference_steps=20,36    controlnet_checkbox=False,        37    controlnet_strength=0.0,          38    controlnet_mode="edge_detection", 39    controlnet_image=None,            40    ip_adapter_checkbox=False,      41    ip_adapter_scale=0.0,   42    ip_adapter_image=None,            43    remove_bg=None,     44    progress=gr.Progress(track_tqdm=True),    45):  46    ckpt_dir='./lora_pussinboots_logos'47    unet_sub_dir = os.path.join(ckpt_dir, "unet")48    #text_encoder_sub_dir = os.path.join(ckpt_dir, "text_encoder")49 50    if model_id is None:51        raise ValueError("Please specify the base model name or path")52 53    generator = torch.Generator(device).manual_seed(seed)54    params = {'prompt': prompt,55              'negative_prompt': negative_prompt,56              'guidance_scale': guidance_scale,57              'num_inference_steps': num_inference_steps,58              'width': width,59              'height': height,60              'generator': generator61             }62 63    if controlnet_checkbox:64        if controlnet_mode == "depth_map":65            controlnet = ControlNetModel.from_pretrained(66                "lllyasviel/sd-controlnet-depth",67                cache_dir="./models_cache",68                torch_dtype=torch_dtype69            )70        elif controlnet_mode == "pose_estimation":71            controlnet = ControlNetModel.from_pretrained(72                "lllyasviel/sd-controlnet-openpose",73                cache_dir="./models_cache",74                torch_dtype=torch_dtype75            )76        elif controlnet_mode == "normal_map":77            controlnet = ControlNetModel.from_pretrained(78                "lllyasviel/sd-controlnet-normal",79                cache_dir="./models_cache",80                torch_dtype=torch_dtype81            )82        elif controlnet_mode == "scribbles":83            controlnet = ControlNetModel.from_pretrained(84                "lllyasviel/sd-controlnet-scribble",85                cache_dir="./models_cache",86                torch_dtype=torch_dtype87            )88        else:89            controlnet = ControlNetModel.from_pretrained(90                "lllyasviel/sd-controlnet-canny",91                cache_dir="./models_cache",92                torch_dtype=torch_dtype93            )94        pipe = StableDiffusionControlNetPipeline.from_pretrained(model_id, 95                                                                 controlnet=controlnet,96                                                                 torch_dtype=torch_dtype, 97                                                                 safety_checker=None).to(device)98        params['image'] = controlnet_image99        params['controlnet_conditioning_scale'] = float(controlnet_strength)100    else:101        pipe = StableDiffusionPipeline.from_pretrained(model_id, 102                                                       torch_dtype=torch_dtype, 103                                                       safety_checker=None).to(device)104 105    pipe.unet = PeftModel.from_pretrained(pipe.unet, unet_sub_dir)106    #pipe.text_encoder = PeftModel.from_pretrained(pipe.text_encoder, text_encoder_sub_dir)107 108    # исправляем ошибку устанорвки lora_scale - меняем на параметр "cross_attention_kwargs"109    # pipe.unet.load_state_dict({k: lora_scale*v for k, v in pipe.unet.state_dict().items()})110    params['cross_attention_kwargs'] = {"scale": float(lora_scale)}111    #pipe.text_encoder.load_state_dict({k: lora_scale*v for k, v in pipe.text_encoder.state_dict().items()})112    113    if torch_dtype in (torch.float16, torch.bfloat16):114        pipe.unet.half()115        #pipe.text_encoder.half()116 117    if ip_adapter_checkbox:118        pipe.load_ip_adapter("h94/IP-Adapter", subfolder="models", weight_name="ip-adapter-plus_sd15.bin")119        pipe.set_ip_adapter_scale(ip_adapter_scale)120        params['ip_adapter_image'] = ip_adapter_image121 122    pipe.to(device)123 124    image = pipe(**params).images[0]125 126    # Если выбрано удаление фона127    if remove_bg:128        image = remove(image)       129 130    return image131 132examples = [133    "Puss in Boots wearing a sombrero crosses the Grand Canyon on a tightrope with a guitar.",134    "Cat wearing a sombrero crosses the Grand Canyon on a tightrope with a guitar.",135    "A cat is playing a song called ""About the Cat"" on an accordion by the sea at sunset. The sun is quickly setting behind the horizon, and the light is fading.",136    "A cat walks through the grass on the streets of an abandoned city. The camera view is always focused on the cat's face.",137    "A young lady in a Russian embroidered kaftan is sitting on a beautiful carved veranda, holding a cup to her mouth and drinking tea from the cup. With her other hand, the girl holds a saucer. The cup and saucer are painted with gzhel. Next to the girl on the table stands a samovar, and steam can be seen above it.",138    "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k",139    "An astronaut riding a green horse",140    "A delicious ceviche cheesecake slice",141]142 143css = """144#col-container {145    margin: 0 auto;146    max-width: 640px;147}148"""149 150def controlnet_params(show_extra):151    return gr.update(visible=show_extra)152    153with gr.Blocks(css=css, fill_height=True) as demo:154    with gr.Column(elem_id="col-container"):155        gr.Markdown(" # Text-to-Image demo")156 157        with gr.Row():158            model_id = gr.Textbox(159                label="Model ID",160                max_lines=1,161                placeholder="Enter model id",162                value=model_id_default,163            )164 165        prompt = gr.Textbox(166            label="Prompt",167            max_lines=1,168            placeholder="Enter your prompt",169        )170        171        negative_prompt = gr.Textbox(172            label="Negative prompt",173            max_lines=1,174            placeholder="Enter your negative prompt",175        )176        177        with gr.Row():178            seed = gr.Number(179                label="Seed",180                minimum=0,181                maximum=MAX_SEED,182                step=1,183                value=42,184            )185            186            guidance_scale = gr.Slider(187                label="Guidance scale",188                minimum=0.0,189                maximum=30.0,190                step=0.1,191                value=7.0,  # Replace with defaults that work for your model192            )193        with gr.Row():194            lora_scale = gr.Slider(195                label="LoRA scale",196                minimum=0.0,197                maximum=1.0,198                step=0.01,199                value=1.0,200            )201 202            num_inference_steps = gr.Slider(203                label="Number of inference steps",204                minimum=1,205                maximum=100,206                step=1,207                value=20,  # Replace with defaults that work for your model208            )209        with gr.Row():210            controlnet_checkbox = gr.Checkbox(211                label="ControlNet",212                value=False213            )214            with gr.Column(visible=False) as controlnet_params:215                controlnet_strength = gr.Slider(216                    label="ControlNet conditioning scale",217                    minimum=0.0,218                    maximum=1.0,219                    step=0.01,220                    value=1.0,  221                )222                controlnet_mode = gr.Dropdown(223                    label="ControlNet mode",224                    choices=["edge_detection", 225                             "depth_map",226                             "pose_estimation", 227                             "normal_map",228                             "scribbles"],229                    value="edge_detection",230                    max_choices=1231                )232                controlnet_image = gr.Image(233                    label="ControlNet condition image",234                    type="pil",235                    format="png"236                )237            controlnet_checkbox.change(238                fn=lambda x: gr.Row.update(visible=x),239                inputs=controlnet_checkbox,240                outputs=controlnet_params241            )242 243        with gr.Row():244            ip_adapter_checkbox = gr.Checkbox(245                label="IPAdapter",246                value=False247            )248            with gr.Column(visible=False) as ip_adapter_params:249                ip_adapter_scale = gr.Slider(250                    label="IPAdapter scale",251                    minimum=0.0,252                    maximum=1.0,253                    step=0.01,254                    value=1.0,  255                )256                ip_adapter_image = gr.Image(257                    label="IPAdapter condition image",258                    type="pil"259                )260            ip_adapter_checkbox.change(261                fn=lambda x: gr.Row.update(visible=x),262                inputs=ip_adapter_checkbox,263                outputs=ip_adapter_params264            )265            266        with gr.Accordion("Optional Settings", open=False):267            268            with gr.Row():269                width = gr.Slider(270                    label="Width",271                    minimum=256,272                    maximum=MAX_IMAGE_SIZE,273                    step=32,274                    value=512,  # Replace with defaults that work for your model275                )276 277                height = gr.Slider(278                    label="Height",279                    minimum=256,280                    maximum=MAX_IMAGE_SIZE,281                    step=32,282                    value=512,  # Replace with defaults that work for your model283                )284 285                # Удаление фона------------------------------------------------------------------------------------------------286                # Checkbox для удаления фона287                remove_bg = gr.Checkbox(288                    label="Remove Background",289                    value=False,290                    interactive=True291                )292                # -------------------------------------------------------------------------------------------------------------293 294        295        gr.Examples(examples=examples, inputs=[prompt])296 297 298        run_button = gr.Button("Run", scale=0, variant="primary")299        result = gr.Image(label="Result", show_label=False)300            301    gr.on(302        triggers=[run_button.click],303        fn=infer,304        inputs=[305            prompt,306            negative_prompt,307            width,308            height,309            model_id,310            seed,311            guidance_scale,      312            lora_scale,313            num_inference_steps,314            controlnet_checkbox,315            controlnet_strength,316            controlnet_mode,317            controlnet_image,318            ip_adapter_checkbox,319            ip_adapter_scale,320            ip_adapter_image, 321            remove_bg, 322        ],323        outputs=[result],324    )325 326if __name__ == "__main__":327    demo.launch()