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 ")]
return text.strip()
def svg_to_html(svg_code: str) -> str:
"""Wrap raw SVG in an HTML container for rendering in Gradio."""
return f'