from __future__ import annotations from pathlib import Path import gradio as gr import numpy as np import torch from model import ConditionalVAE from PIL import Image from safetensors.torch import load_file ARTIFACT = Path(__file__).resolve().parent / "artifacts" / "glyph-forge-cvae" MODEL = ConditionalVAE() MODEL.load_state_dict(load_file(ARTIFACT / "model.safetensors")) MODEL.eval() def generate_digit(label: int, seed: int) -> Image.Image: generator = torch.Generator().manual_seed(seed) latent = torch.randn(1, MODEL.latent_dimensions, generator=generator) with torch.no_grad(): pixels = MODEL.decode(latent, torch.tensor([label]))[0] image = np.clip(pixels.reshape(8, 8).numpy() * 255, 0, 255).astype(np.uint8) return Image.fromarray(image, mode="L").resize((512, 512), Image.Resampling.NEAREST) with gr.Blocks(title="GlyphForge CVAE") as demo: gr.Markdown("# GlyphForge\nSample the eight-dimensional latent style of any digit.") with gr.Row(): label = gr.Slider(0, 9, value=3, step=1, label="Digit") seed = gr.Slider(0, 100_000, value=2031, step=1, label="Latent seed") output = gr.Image(value=generate_digit(3, 2031), label="Generated 8x8 glyph") button = gr.Button("Forge glyph", variant="primary") button.click(generate_digit, inputs=[label, seed], outputs=output) if __name__ == "__main__": demo.launch()