virtual-tryon / app.py
Nandha2017's picture
Upload folder using huggingface_hub
d04b5a9 verified
Raw History Blame
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
# ---------------------------------------------------------------------------
@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()