macrdel/text2img_diff_model
0
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 