from __future__ import annotations from pathlib import Path import gradio as gr import torch from model import ConditionalEnergyNetwork, langevin_sample from PIL import Image from safetensors.torch import load_file ARTIFACT_DIR = Path(__file__).resolve().parent / "artifacts" / "energy-pocket" MODEL = ConditionalEnergyNetwork() MODEL.load_state_dict(load_file(ARTIFACT_DIR / "model.safetensors")) MODEL.eval() def sample_energy( label: int, seed: int, steps: int, ) -> tuple[Image.Image, dict]: generator = torch.Generator().manual_seed(int(seed)) labels = torch.full((12,), int(label), dtype=torch.long) initial = torch.rand(12, 64, generator=generator) generated = langevin_sample( MODEL, initial, labels, steps=int(steps), step_size=0.08, noise_scale=0.008, generator=generator, ) canvas = Image.new("L", (512, 384), color=0) for index, pixels in enumerate(generated): image = Image.fromarray( pixels.reshape(8, 8).mul(255).to(torch.uint8).numpy(), mode="L", ).resize((120, 120), Image.Resampling.NEAREST) canvas.paste(image, ((index % 4) * 128 + 4, (index // 4) * 128 + 4)) with torch.inference_mode(): energies = MODEL(generated, labels) return canvas, { "digit": int(label), "langevin_steps": int(steps), "mean_energy": float(energies.mean()), "energy_standard_deviation": float(energies.std()), } with gr.Blocks(title="Energy Pocket") as demo: gr.Markdown( "# Energy Pocket\n" "Watch random pixels descend a learned class-conditional energy landscape " "through stochastic Langevin dynamics." ) with gr.Row(): label = gr.Slider(0, 9, value=2, step=1, label="Digit class") seed = gr.Slider(0, 100_000, value=2069, step=1, label="Seed") steps = gr.Slider(10, 160, value=40, step=10, label="Langevin steps") initial = sample_energy(2, 2069, 40) output = gr.Image(value=initial[0], label="Energy-minimized samples") metrics = gr.JSON(value=initial[1], label="Energy readout") button = gr.Button("Descend the energy landscape", variant="primary") button.click(sample_energy, inputs=[label, seed, steps], outputs=[output, metrics]) if __name__ == "__main__": demo.launch()