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