""" 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 # --------------------------------------------------------------------------- @spaces.GPU(duration=120) 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()