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 block.""" if "" in text: text = text[: text.find("") + len("")] return text.strip() def svg_to_html(svg_code: str) -> str: """Wrap raw SVG in an HTML container for rendering in Gradio.""" return f'
{svg_code}
' @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"""
VFig logo

VFig Vectorizing Complex Figures in SVG with Vision-Language Models

""") 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())