zixianma02's picture
Update app.py
cc2ac41
Raw History Blame Contribute Delete
6.5 kB
import argparse
import base64
import os
import tempfile
import spaces # must be imported before torch/CUDA on ZeroGPU Spaces
import gradio as gr
import torch
from huggingface_hub import snapshot_download
from PIL import Image
from transformers import AutoModelForImageTextToText, AutoProcessor
# -------------------------------------------------------------
# ARGS
# -------------------------------------------------------------
parser = argparse.ArgumentParser()
parser.add_argument(
"--base_model",
type=str,
default="XunmeiLiu/VFIG-4B",
help="HuggingFace repo ID or local path to the base (or merged) model",
)
args = parser.parse_args()
# -------------------------------------------------------------
# CONFIG
# -------------------------------------------------------------
PROMPT_TEXT = "Convert this figure into valid SVG code."
MAX_NEW_TOKENS = 8192
# -------------------------------------------------------------
# LOAD MODEL (to CPU — GPU is only available inside @spaces.GPU functions)
# -------------------------------------------------------------
# Use persistent storage cache if available, otherwise fall back to HF default cache
PERSISTENT_CACHE = "/data/models"
if os.path.exists("/data"):
model_cache_path = os.path.join(PERSISTENT_CACHE, args.base_model.replace("/", "--"))
else:
model_cache_path = None
if model_cache_path and os.path.exists(model_cache_path):
print(f"Loading model from persistent cache: {model_cache_path}")
args.base_model = model_cache_path
elif not os.path.exists(args.base_model):
print(f"Downloading base model from HuggingFace Hub: {args.base_model} …")
if model_cache_path:
os.makedirs(PERSISTENT_CACHE, exist_ok=True)
args.base_model = snapshot_download(repo_id=args.base_model, local_dir=model_cache_path)
else:
args.base_model = snapshot_download(repo_id=args.base_model)
print(f"Model downloaded to: {args.base_model}")
print("Loading processor …")
processor = AutoProcessor.from_pretrained(args.base_model, trust_remote_code=True)
print("Loading model …")
model = AutoModelForImageTextToText.from_pretrained(
args.base_model,
torch_dtype=torch.bfloat16,
trust_remote_code=True,
)
model.eval()
print("Model ready.")
# -------------------------------------------------------------
# HELPERS
# -------------------------------------------------------------
def extract_svg(text: str) -> str:
"""Keep only the <svg … </svg> block."""
if "<svg" in text:
text = text[text.find("<svg"):]
if "</svg>" in text:
text = text[: text.find("</svg>") + len("</svg>")]
return text.strip()
def svg_to_html(svg_code: str) -> str:
"""Wrap raw SVG in an HTML container for rendering in Gradio."""
return f'<div style="display:flex;justify-content:center;align-items:center;min-height:300px;background:#fff;border:1px solid #ddd;border-radius:8px;padding:16px;">{svg_code}</div>'
@spaces.GPU(duration=120)
def predict(input_image: Image.Image):
"""Take a PIL image, return (svg_code, rendered_html, svg_file_path)."""
if input_image is None:
return "", "", None
device = "cuda" if torch.cuda.is_available() else "cpu"
model.to(device)
img = input_image.convert("RGB")
messages = [
{
"role": "user",
"content": [
{"type": "image"},
{"type": "text", "text": PROMPT_TEXT},
],
}
]
chat_input = processor.tokenizer.apply_chat_template(
messages, tokenize=False, add_generation_prompt=True
)
model_inputs = processor(
text=[chat_input], images=[img], return_tensors="pt"
).to(device)
with torch.no_grad():
output_ids = model.generate(
**model_inputs,
max_new_tokens=MAX_NEW_TOKENS,
do_sample=False,
)
decoded = processor.tokenizer.decode(output_ids[0], skip_special_tokens=True)
svg_code = extract_svg(decoded)
# Render the SVG inline via HTML
rendered_html = svg_to_html(svg_code)
# Write to a temp file so the user can download it
tmp = tempfile.NamedTemporaryFile(suffix=".svg", delete=False, mode="w")
tmp.write(svg_code)
tmp.close()
return svg_code, rendered_html, tmp.name
# -------------------------------------------------------------
# GRADIO UI
# -------------------------------------------------------------
LOGO_PATH = os.path.join(os.path.dirname(os.path.abspath(__file__)), "fig.png")
with open(LOGO_PATH, "rb") as f:
logo_b64 = base64.b64encode(f.read()).decode()
CSS = """
.header-row {
display: flex;
align-items: center;
justify-content: center;
gap: 16px;
padding: 24px 0 8px 0;
}
.header-row img {
height: 64px;
border-radius: 10px;
}
.header-row h1 {
margin: 0;
font-size: 2.4em;
white-space: nowrap;
}
.header-row h1 span {
font-weight: 400;
opacity: 0.65;
font-size: 0.55em;
margin-left: 12px;
}
footer { display: none !important; }
#right-col { gap: 0 !important; }
.gr-samples-table td img { height: 900px !important; width: 1600px !important; object-fit: contain !important; }
.gr-samples-table td { padding: 8px !important; }
"""
with gr.Blocks(title="VFig") as demo:
gr.HTML(f"""
<div class="header-row">
<img src="data:image/png;base64,{logo_b64}" alt="VFig logo">
<h1>VFig <span>Vectorizing Complex Figures in SVG with Vision-Language Models</span></h1>
</div>
""")
with gr.Row():
with gr.Column(scale=1):
input_image = gr.Image(type="pil", label="Input Image")
run_btn = gr.Button("Generate SVG", variant="primary", size="lg")
with gr.Column(scale=1, elem_id="right-col"):
with gr.Tabs():
with gr.Tab("Rendered SVG"):
rendered_preview = gr.HTML(container=False)
with gr.Tab("SVG Code"):
svg_code = gr.Code(language="html", lines=20)
svg_file = gr.File(label="Download SVG")
gr.Examples(
examples=[
["./simple_diagram.png"],
["./medium_diagram.png"],
["./complex_diagram.png"],
],
inputs=[input_image],
label="Examples",
examples_per_page=3,
)
run_btn.click(
fn=predict,
inputs=[input_image],
outputs=[svg_code, rendered_preview, svg_file],
)
demo.launch(css=CSS, theme=gr.themes.Soft())