Spaces:
Running on Zero
Running on Zero
| from __future__ import annotations | |
| import os | |
| import shutil | |
| import tempfile | |
| from pathlib import Path | |
| from types import SimpleNamespace | |
| import gradio as gr | |
| import spaces | |
| import torch | |
| from diffusers import DDIMScheduler | |
| from diffusers.models import AutoencoderKL | |
| from huggingface_hub import snapshot_download | |
| from PIL import Image, ImageChops | |
| from removal.v1_2 import RemovalSDXLPipeline, build_removal_model, load_cfg, load_removal_model | |
| MODEL_ID = "hustvl/Moebius" | |
| VAE_ID = "hustvl/PixelHacker" | |
| LOCAL_MODEL_DIR = Path("weight/Moebius") | |
| LOCAL_VAE_DIR = Path("weight/vae") | |
| MODEL_CONFIG = "config/model_cfg/moebius.yaml" | |
| CHECKPOINTS = { | |
| "Fine-tuned CelebA-HQ": "ft_celebahq", | |
| "Fine-tuned FFHQ": "ft_ffhq", | |
| "Fine-tuned Places2": "ft_places2", | |
| "Pretrained": "pretrained", | |
| } | |
| def _resize_for_model(image: Image.Image, mask: Image.Image, image_size: int): | |
| target_size = (int(image_size), int(image_size)) | |
| return ( | |
| image.resize(target_size, Image.Resampling.LANCZOS), | |
| mask.resize(target_size, Image.Resampling.NEAREST), | |
| ) | |
| def _ensure_weights() -> None: | |
| moebius_path = Path( | |
| snapshot_download( | |
| repo_id=MODEL_ID, | |
| allow_patterns=[ | |
| "pretrained/diffusion_pytorch_model.bin", | |
| "ft_celebahq/diffusion_pytorch_model.bin", | |
| "ft_ffhq/diffusion_pytorch_model.bin", | |
| "ft_places2/diffusion_pytorch_model.bin", | |
| ], | |
| ) | |
| ) | |
| vae_path = Path(snapshot_download(repo_id=VAE_ID, allow_patterns=["vae/*"])) | |
| LOCAL_MODEL_DIR.mkdir(parents=True, exist_ok=True) | |
| LOCAL_VAE_DIR.mkdir(parents=True, exist_ok=True) | |
| for checkpoint in CHECKPOINTS.values(): | |
| target_dir = LOCAL_MODEL_DIR / checkpoint | |
| target_dir.mkdir(parents=True, exist_ok=True) | |
| shutil.copy2( | |
| moebius_path / checkpoint / "diffusion_pytorch_model.bin", | |
| target_dir / "diffusion_pytorch_model.bin", | |
| ) | |
| shutil.copy2(vae_path / "vae" / "config.json", LOCAL_VAE_DIR / "config.json") | |
| shutil.copy2( | |
| vae_path / "vae" / "diffusion_pytorch_model.bin", | |
| LOCAL_VAE_DIR / "diffusion_pytorch_model.bin", | |
| ) | |
| def _build_pipeline(checkpoint_key: str) -> RemovalSDXLPipeline: | |
| model_cfg = load_cfg(MODEL_CONFIG) | |
| weight_path = LOCAL_MODEL_DIR / checkpoint_key / "diffusion_pytorch_model.bin" | |
| removal_model = build_removal_model(model_cfg, 20) | |
| load_removal_model(removal_model, str(weight_path), device="cuda") | |
| vae = AutoencoderKL.from_pretrained(str(LOCAL_VAE_DIR)) | |
| scheduler = DDIMScheduler( | |
| beta_start=0.00085, | |
| beta_end=0.012, | |
| beta_schedule="scaled_linear", | |
| num_train_timesteps=1000, | |
| clip_sample=False, | |
| ) | |
| return RemovalSDXLPipeline( | |
| removal_model=removal_model, | |
| vae=vae, | |
| scheduler=scheduler, | |
| device="cuda", | |
| dtype=torch.float, | |
| ) | |
| _ensure_weights() | |
| PIPELINES = {key: _build_pipeline(value) for key, value in CHECKPOINTS.items()} | |
| def _editor_to_image_and_mask(editor_value): | |
| if editor_value is None: | |
| raise gr.Error("Upload an image and paint a mask first.") | |
| background = editor_value.get("background") | |
| layers = editor_value.get("layers") or [] | |
| if background is None: | |
| raise gr.Error("Upload an image first.") | |
| if not layers: | |
| raise gr.Error("Paint a mask over the area to inpaint.") | |
| image = background.convert("RGB") | |
| mask = Image.new("L", image.size, 0) | |
| for layer in layers: | |
| if layer is None: | |
| continue | |
| alpha = layer.convert("RGBA").getchannel("A") | |
| mask = ImageChops.lighter(mask, alpha) | |
| if mask.getbbox() is None: | |
| raise gr.Error("Paint a visible mask over the area to inpaint.") | |
| return image, mask | |
| def inpaint( | |
| editor_value, | |
| checkpoint_label, | |
| image_size, | |
| steps, | |
| guidance_scale, | |
| mask_dilate, | |
| paste, | |
| compensate, | |
| noise_offset, | |
| ): | |
| image, mask = _editor_to_image_and_mask(editor_value) | |
| image, mask = _resize_for_model(image, mask, int(image_size)) | |
| pipe = PIPELINES[checkpoint_label] | |
| result = pipe( | |
| [image], | |
| [mask], | |
| image_size=int(image_size), | |
| mask_dilate_kernel_size=int(mask_dilate), | |
| mask_preprocess_type="dilate", | |
| num_steps=int(steps), | |
| guidance_scale=float(guidance_scale), | |
| paste=bool(paste), | |
| compensate=bool(compensate), | |
| noise_offset=float(noise_offset), | |
| mute=True, | |
| )[0] | |
| with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as output: | |
| result.save(output.name) | |
| return output.name | |
| with gr.Blocks(title="Moebius Inpainting") as demo: | |
| gr.Markdown("# Moebius Inpainting") | |
| gr.Markdown(f"[Model on Hugging Face](https://huggingface.co/{MODEL_ID})") | |
| with gr.Row(): | |
| with gr.Column(scale=2): | |
| editor = gr.ImageEditor( | |
| label="Image and mask", | |
| type="pil", | |
| sources=["upload"], | |
| brush=gr.Brush(colors=["#ffffff"], color_mode="fixed"), | |
| height=520, | |
| ) | |
| run_button = gr.Button("Inpaint", variant="primary") | |
| with gr.Column(scale=1): | |
| checkpoint = gr.Dropdown( | |
| label="Checkpoint", | |
| choices=list(CHECKPOINTS.keys()), | |
| value="Fine-tuned Places2", | |
| ) | |
| image_size = gr.Radio( | |
| label="Resolution", | |
| choices=[512, 768], | |
| value=512, | |
| ) | |
| steps = gr.Slider(5, 30, value=20, step=1, label="Steps") | |
| guidance = gr.Slider(1.0, 6.0, value=2.0, step=0.1, label="CFG") | |
| mask_dilate = gr.Slider(0, 64, value=0, step=1, label="Mask dilation") | |
| noise_offset = gr.Slider(0.0, 0.1, value=0.0357, step=0.0001, label="Noise offset") | |
| paste = gr.Checkbox(label="Paste into original image", value=True) | |
| compensate = gr.Checkbox(label="Color compensation", value=False) | |
| output = gr.Image(label="Inpainted image", type="filepath") | |
| run_button.click( | |
| fn=inpaint, | |
| inputs=[ | |
| editor, | |
| checkpoint, | |
| image_size, | |
| steps, | |
| guidance, | |
| mask_dilate, | |
| paste, | |
| compensate, | |
| noise_offset, | |
| ], | |
| outputs=output, | |
| ) | |
| gr.Markdown("[Twitter / X](https://x.com/realmrfakename)") | |
| if __name__ == "__main__": | |
| demo.queue().launch() | |