Spaces:
Runtime error
Runtime error
| import gradio as gr | |
| import torch | |
| from diffusers import StableDiffusionPipeline, DPMSolverMultistepScheduler, UNet2DConditionModel | |
| from peft import PeftModel | |
| MODEL_ID = "Manojb/stable-diffusion-2-1-base" | |
| LORA_PATHS = { | |
| "akiec": "./lora/lora_akiec_final", | |
| "bcc": "./lora/lora_bcc_final", | |
| "df": "./lora/lora_df_final", | |
| "mel": "./lora/lora_mel_final", | |
| "vasc": "./lora/lora_vasc_final", | |
| } | |
| CLASS_NEGATIVE_PROMPTS = { | |
| "mel": "benign nevus, symmetric, uniform color, regular border, blurry, low quality, artifacts, text, watermark, cartoon, non-dermoscopic, overexposed", | |
| "bcc": "melanoma, nevus, pigment network, blurry, low quality, artifacts, text, watermark, cartoon, non-dermoscopic, overexposed", | |
| "akiec": "melanoma, smooth surface, no scale, blurry, low quality, artifacts, text, watermark, cartoon, non-dermoscopic, overexposed", | |
| "df": "melanoma, irregular border, blue-white veil, blurry, low quality, artifacts, text, watermark, cartoon, non-dermoscopic, overexposed", | |
| "vasc": "melanoma, pigment network, brown color, blurry, low quality, artifacts, text, watermark, cartoon, non-dermoscopic, overexposed", | |
| } | |
| # Load base pipeline | |
| pipe = StableDiffusionPipeline.from_pretrained( | |
| MODEL_ID, | |
| torch_dtype=torch.float32, | |
| safety_checker=None, | |
| requires_safety_checker=False, | |
| ) | |
| pipe.scheduler = DPMSolverMultistepScheduler.from_config(pipe.scheduler.config) | |
| pipe.enable_attention_slicing() | |
| _current_lora = None | |
| _base_unet_state = None | |
| def apply_lora(name): | |
| global _current_lora, pipe | |
| if _current_lora == name: | |
| return | |
| # Reload fresh UNet then apply LoRA | |
| pipe.unet = UNet2DConditionModel.from_pretrained( | |
| MODEL_ID, subfolder="unet", torch_dtype=torch.float32 | |
| ) | |
| pipe.unet = PeftModel.from_pretrained(pipe.unet, LORA_PATHS[name]) | |
| _current_lora = name | |
| def generate(prompt, lora_name, steps): | |
| apply_lora(lora_name) | |
| image = pipe( | |
| prompt, | |
| negative_prompt=CLASS_NEGATIVE_PROMPTS[lora_name], | |
| num_inference_steps=int(steps), | |
| guidance_scale=9.0, | |
| ).images[0] | |
| return image | |
| with gr.Blocks() as demo: | |
| gr.Markdown("# 🧠 DermaDiff — SD 2.1 Generator") | |
| gr.Markdown("Generate dermoscopy images based on disease class using LoRA") | |
| with gr.Row(): | |
| with gr.Column(): | |
| prompt = gr.Textbox(label="Prompt", value="dermoscopy image of melanoma") | |
| lora = gr.Dropdown(choices=list(LORA_PATHS.keys()), value="mel", label="Disease Class") | |
| steps = gr.Slider(10, 50, value=50, step=1, label="Steps") | |
| btn = gr.Button("Generate", variant="primary") | |
| with gr.Column(): | |
| output = gr.Image(label="Generated Image") | |
| btn.click(fn=generate, inputs=[prompt, lora, steps], outputs=output) | |
| demo.launch() |