Spaces:
Running on Zero
Running on Zero
File size: 7,557 Bytes
bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 d04b5a9 bb42d68 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 | """
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()
|