ARotting's picture
Publish Conditional energy model with persistent contrastive divergence
7d90be6 verified
Raw
History Blame Contribute Delete
2.36 kB
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()