Team Ai
Apppublic

macrdel/text2img_diff_model

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
app.py228 linesDownload Raw Back to root
1import gradio as gr2import numpy as np3import random4 5from diffusers import StableDiffusionPipeline # DiffusionPipeline6from peft import PeftModel, PeftConfig7import torch8 9device = "cuda" if torch.cuda.is_available() else "cpu"10 11# Model list including your LoRA model12MODEL_LIST = [13    "CompVis/stable-diffusion-v1-4",14    "stabilityai/sdxl-turbo",15    "runwayml/stable-diffusion-v1-5",16    "stabilityai/stable-diffusion-2-1",17    "macrdel/unico_proj",18]19 20if torch.cuda.is_available():21    torch_dtype = torch.float1622else:23    torch_dtype = torch.float3224 25# Cache to avoid re-initializing pipelines repeatedly26model_cache = {}27 28def load_pipeline(model_id: str, lora_scale):29    """30    Loads or retrieves a cached DiffusionPipeline.31    32    If the chosen model is your LoRA adapter, then load the base model 33    (CompVis/stable-diffusion-v1-4) and apply the LoRA weights.34    """35    if model_id in model_cache:36        return model_cache[model_id]37    38    if model_id == "macrdel/unico_proj":39        # Use the specified base model for your LoRA adapter.40        base_model = "CompVis/stable-diffusion-v1-4"41        pipe = StableDiffusionPipeline.from_pretrained(base_model, torch_dtype=torch_dtype)42        # Load the LoRA weights43        pipe.unet = PeftModel.from_pretrained(44            pipe.unet, 45            model_id, 46            subfolder="unet", 47            torch_dtype=torch_dtype48        )49        pipe.text_encoder = PeftModel.from_pretrained(50            pipe.text_encoder, 51            model_id, 52            subfolder="text_encoder", 53            torch_dtype=torch_dtype54        )55    else:56        pipe = StableDiffusionPipeline.from_pretrained(model_id, torch_dtype=torch_dtype, safety_checker=None).to(device)57        pipe.unet.load_state_dict({k: lora_scale * v for k, v in pipe.unet.state_dict().items()})58        pipe.text_encoder.load_state_dict({k: lora_scale * v for k, v in pipe.text_encoder.state_dict().items()})59    60    pipe.to(device)61    model_cache[model_id] = pipe62    return pipe63 64MAX_SEED = np.iinfo(np.int32).max65MAX_IMAGE_SIZE = 102466 67def infer(68    model_id,69    prompt,70    negative_prompt,71    seed,72    randomize_seed,73    width,74    height,75    guidance_scale,76    num_inference_steps,77    lora_scale,  # New parameter for adjusting LoRA scale78    progress=gr.Progress(track_tqdm=True),79):80    # Load the pipeline for the chosen model81    pipe = load_pipeline(model_id, lora_scale)82 83    if randomize_seed:84        seed = random.randint(0, MAX_SEED)85 86    generator = torch.Generator(device=device).manual_seed(seed)87 88    # If using the LoRA model, update the LoRA scale if supported.89    # if model_id == "macrdel/unico_proj":90        # This assumes your pipeline's unet has a method to update the LoRA scale.91        # if hasattr(pipe.unet, "set_lora_scale"):92        #    pipe.unet.set_lora_scale(lora_scale)93        # else:94        #    print("Warning: LoRA scale adjustment method not found on UNet.")95 96    image = pipe(97        prompt=prompt,98        negative_prompt=negative_prompt,99        guidance_scale=guidance_scale,100        num_inference_steps=num_inference_steps,101        width=width,102        height=height,103        generator=generator,104    ).images[0]105 106    return image, seed107 108examples = [109    "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k",110    "An astronaut riding a green horse",111    "A delicious ceviche cheesecake slice",112]113 114css = """115#col-container {116    margin: 0 auto;117    max-width: 640px;118}119"""120 121with gr.Blocks(css=css) as demo:122    with gr.Column(elem_id="col-container"):123        gr.Markdown(" # Text-to-Image Gradio Template")124 125        with gr.Row():126            # Dropdown to select the model from Hugging Face127            model_id = gr.Dropdown(128                label="Model",129                choices=MODEL_LIST,130                value=MODEL_LIST[0],  # Default model131            )132 133        with gr.Row():134            prompt = gr.Text(135                label="Prompt",136                show_label=False,137                max_lines=1,138                placeholder="Enter your prompt",139                container=False,140            )141 142            run_button = gr.Button("Run", scale=0, variant="primary")143 144        result = gr.Image(label="Result", show_label=False)145 146        with gr.Accordion("Advanced Settings", open=False):147            negative_prompt = gr.Text(148                label="Negative prompt",149                max_lines=1,150                placeholder="Enter a negative prompt",151            )152 153            seed = gr.Slider(154                label="Seed",155                minimum=0,156                maximum=MAX_SEED,157                step=1,158                value=42,  # Default seed159            )160 161            randomize_seed = gr.Checkbox(label="Randomize seed", value=True)162 163            with gr.Row():164                width = gr.Slider(165                    label="Width",166                    minimum=256,167                    maximum=MAX_IMAGE_SIZE,168                    step=32,169                    value=1024,170                )171 172                height = gr.Slider(173                    label="Height",174                    minimum=256,175                    maximum=MAX_IMAGE_SIZE,176                    step=32,177                    value=1024,178                )179 180            with gr.Row():181                guidance_scale = gr.Slider(182                    label="Guidance scale",183                    minimum=0.0,184                    maximum=20.0,185                    step=0.5,186                    value=7.0,187                )188 189                num_inference_steps = gr.Slider(190                    label="Number of inference steps",191                    minimum=1,192                    maximum=100,193                    step=1,194                    value=20,195                )196 197            # New slider for LoRA scale.198            lora_scale = gr.Slider(199                label="LoRA Scale",200                minimum=0.0,201                maximum=2.0,202                step=0.1,203                value=1.0,204                info="Adjust the influence of the LoRA weights",205            )206 207        gr.Examples(examples=examples, inputs=[prompt])208    gr.on(209        triggers=[run_button.click, prompt.submit],210        fn=infer,211        inputs=[212            model_id,213            prompt,214            negative_prompt,215            seed,216            randomize_seed,217            width,218            height,219            guidance_scale,220            num_inference_steps,221            lora_scale,  # Pass the new slider value222        ],223        outputs=[result, seed],224    )225 226if __name__ == "__main__":227    demo.launch()228