Spaces:
Running on Zero
Running on Zero
Download app.py from Nandha2017/virtual-tryon: direct link, hf CLI and curl.
- Browser
- Download file 7.56 kB
-
https://huggingface.co/spaces/Nandha2017/virtual-tryon/resolve/bd4b4565afbb0483c196d3a77384b50ef60b1723/app.py
- Command line
-
hf download hf://spaces/Nandha2017/virtual-tryon@bd4b4565afbb0483c196d3a77384b50ef60b1723/app.py
-
curl -L -o app.py https://huggingface.co/spaces/Nandha2017/virtual-tryon/resolve/bd4b4565afbb0483c196d3a77384b50ef60b1723/app.py
7.56 kB
| """ | |
| Virtual Try-On — CatVTON + Hugging Face ZeroGPU | |
| No local GPU or model storage needed. Generated images download to your device. | |
| """ | |
| import datetime | |
| import os | |
| import sys | |
| import gradio as gr | |
| import numpy as np | |
| import spaces | |
| import torch | |
| from huggingface_hub import snapshot_download | |
| from PIL import Image, ImageDraw | |
| # --------------------------------------------------------------------------- | |
| # Persistent storage (/data on ZeroGPU Spaces, /tmp fallback) | |
| # --------------------------------------------------------------------------- | |
| DATA_DIR = "/data" if os.path.exists("/data") else "/tmp" | |
| MODELS_DIR = os.path.join(DATA_DIR, "catvton_models") | |
| OUTPUT_DIR = os.path.join(DATA_DIR, "outputs") | |
| os.makedirs(MODELS_DIR, exist_ok=True) | |
| os.makedirs(OUTPUT_DIR, exist_ok=True) | |
| os.environ["HF_HOME"] = os.path.join(DATA_DIR, "hf_cache") | |
| os.environ["HUGGINGFACE_HUB_CACHE"] = os.path.join(DATA_DIR, "hf_cache", "hub") | |
| # --------------------------------------------------------------------------- | |
| # Model download — runs once at Space startup on HF servers (not locally) | |
| # --------------------------------------------------------------------------- | |
| CATVTON_REPO = "zhengchong/CatVTON" | |
| CATVTON_LOCAL = os.path.join(MODELS_DIR, "CatVTON") | |
| def download_models(): | |
| if os.path.exists(os.path.join(CATVTON_LOCAL, "model_index.json")): | |
| print("CatVTON already cached.") | |
| return | |
| print("Downloading CatVTON (~4 GB) to HF persistent storage…") | |
| snapshot_download( | |
| repo_id=CATVTON_REPO, | |
| local_dir=CATVTON_LOCAL, | |
| local_dir_use_symlinks=False, | |
| ignore_patterns=["*.md", "*.txt", "*.py"], | |
| ) | |
| print("CatVTON ready.") | |
| # --------------------------------------------------------------------------- | |
| # Pipeline (loaded lazily inside @spaces.GPU) | |
| # --------------------------------------------------------------------------- | |
| _pipe = None | |
| def _get_pipe(): | |
| global _pipe | |
| if _pipe is not None: | |
| return _pipe | |
| from diffusers import StableDiffusionInpaintPipeline | |
| _pipe = StableDiffusionInpaintPipeline.from_pretrained( | |
| CATVTON_LOCAL, | |
| torch_dtype=torch.float16, | |
| safety_checker=None, | |
| requires_safety_checker=False, | |
| ).to("cuda") | |
| _pipe.set_progress_bar_config(disable=True) | |
| print("Pipeline loaded on CUDA.") | |
| return _pipe | |
| # --------------------------------------------------------------------------- | |
| # Image helpers | |
| # --------------------------------------------------------------------------- | |
| TARGET_SIZE = 512 | |
| def _fit_to_square(img: Image.Image, size: int = TARGET_SIZE) -> Image.Image: | |
| img = img.convert("RGB") | |
| img.thumbnail((size, size), Image.LANCZOS) | |
| canvas = Image.new("RGB", (size, size), (255, 255, 255)) | |
| canvas.paste(img, ((size - img.width) // 2, (size - img.height) // 2)) | |
| return canvas | |
| def _make_mask(size: int, cloth_type: str) -> Image.Image: | |
| mask = Image.new("L", (size, size), 0) | |
| d = ImageDraw.Draw(mask) | |
| if cloth_type == "upper": | |
| d.rectangle([int(size*.10), int(size*.18), int(size*.90), int(size*.65)], fill=255) | |
| elif cloth_type == "lower": | |
| d.rectangle([int(size*.05), int(size*.55), int(size*.95), int(size*1.0)], fill=255) | |
| else: # overall / dress | |
| d.rectangle([int(size*.05), int(size*.15), int(size*.95), int(size*1.0)], fill=255) | |
| return mask | |
| # --------------------------------------------------------------------------- | |
| # ZeroGPU inference | |
| # --------------------------------------------------------------------------- | |
| def run_tryon( | |
| person_image: Image.Image, | |
| garment_image: Image.Image, | |
| cloth_type: str, | |
| num_steps: int, | |
| guidance_scale: float, | |
| seed: int, | |
| ) -> tuple: | |
| if person_image is None or garment_image is None: | |
| raise gr.Error("Please upload both a person photo and a garment image.") | |
| pipe = _get_pipe() | |
| person = _fit_to_square(person_image) | |
| garment = _fit_to_square(garment_image) | |
| mask = _make_mask(TARGET_SIZE, cloth_type) | |
| rng = torch.Generator(device="cuda") | |
| rng.manual_seed(int(seed) if seed != -1 else torch.randint(0, 2**32, (1,)).item()) | |
| prompt = ( | |
| "a person wearing the garment in the reference image, " | |
| "photorealistic, high quality, natural lighting" | |
| ) | |
| negative = "blurry, distorted, deformed, low quality, artifacts" | |
| result = pipe( | |
| prompt=prompt, | |
| negative_prompt=negative, | |
| image=person, | |
| mask_image=mask, | |
| num_inference_steps=num_steps, | |
| guidance_scale=guidance_scale, | |
| generator=rng, | |
| ) | |
| output_images = result.images | |
| timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S") | |
| saved_paths = [] | |
| for i, img in enumerate(output_images): | |
| path = os.path.join(OUTPUT_DIR, f"tryon_{timestamp}_{i}.png") | |
| img.save(path, format="PNG") | |
| saved_paths.append(path) | |
| return output_images, saved_paths | |
| # --------------------------------------------------------------------------- | |
| # Gradio UI | |
| # --------------------------------------------------------------------------- | |
| with gr.Blocks(title="Virtual Try-On", theme=gr.themes.Soft()) as demo: | |
| gr.Markdown( | |
| "# 👗 Virtual Try-On\n" | |
| "Upload a **person photo** and a **garment image**, select the type, then click **Try On**.\n\n" | |
| "> Runs entirely on **Hugging Face ZeroGPU** (free A10G) — no local GPU needed. \n" | |
| "> Models download once to HF persistent storage. Images save to your device via the Download button." | |
| ) | |
| with gr.Row(): | |
| with gr.Column(): | |
| person_input = gr.Image(label="Person Photo", type="pil", height=380) | |
| garment_input = gr.Image(label="Garment Image", type="pil", height=380) | |
| cloth_type = gr.Radio( | |
| ["upper", "lower", "overall"], | |
| value="upper", | |
| label="Garment Type", | |
| info="upper=top/shirt | lower=pants/skirt | overall=dress/full outfit", | |
| ) | |
| with gr.Accordion("Advanced", open=False): | |
| num_steps = gr.Slider(10, 50, value=30, step=1, label="Steps") | |
| guidance = gr.Slider(1.0, 10.0, value=7.5, step=0.5, label="Guidance Scale") | |
| seed_input = gr.Number(label="Seed (-1 = random)", value=-1, precision=0) | |
| try_btn = gr.Button("👗 Try On", variant="primary", size="lg") | |
| with gr.Column(): | |
| output_gallery = gr.Gallery(label="Result", columns=1, height=380) | |
| output_files = gr.File( | |
| label="⬇ Download to your device", | |
| file_count="multiple", | |
| interactive=False, | |
| ) | |
| try_btn.click( | |
| fn=run_tryon, | |
| inputs=[person_input, garment_input, cloth_type, num_steps, guidance, seed_input], | |
| outputs=[output_gallery, output_files], | |
| ) | |
| gr.Markdown( | |
| "---\n" | |
| "**Tips:** front-facing photo · garment on white/neutral background · upper body for shirts\n\n" | |
| "First run: ~2-5 min (model download). Subsequent runs: ~15-30s.\n\n" | |
| "Built with [CatVTON](https://github.com/zhengchong/CatVTON) · " | |
| "[Gradio](https://gradio.app) · [ZeroGPU](https://huggingface.co/docs/hub/spaces-zerogpu)" | |
| ) | |
| download_models() | |
| if __name__ == "__main__": | |
| demo.launch() | |