Spaces:
Running on Zero
Running on Zero
Download app.py from allenai/VFig-Image2SVG-Demo: direct link, hf CLI and curl.
- Browser
- Download file 6.5 kB
-
https://huggingface.co/spaces/allenai/VFig-Image2SVG-Demo/resolve/main/app.py
- Command line
-
hf download hf://spaces/allenai/VFig-Image2SVG-Demo/app.py
-
curl -L -o app.py https://huggingface.co/spaces/allenai/VFig-Image2SVG-Demo/resolve/main/app.py
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>' | |
| 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()) | |