pocket-wgan / source /app.py
ARotting's picture
Publish Projection-conditioned WGAN-GP with collapse diagnostics
f3a1fe5 verified
Raw
History Blame Contribute Delete
2.53 kB
from __future__ import annotations
from pathlib import Path
import gradio as gr
import torch
from model import ConditionalGenerator
from PIL import Image, ImageDraw
from safetensors.torch import load_file
ARTIFACT_DIR = Path(__file__).resolve().parent / "artifacts" / "pocket-wgan"
MODEL = ConditionalGenerator()
MODEL.load_state_dict(load_file(ARTIFACT_DIR / "generator.safetensors"))
MODEL.eval()
def generate_gallery(label: int, seed: int, temperature: float) -> tuple[Image.Image, dict]:
labels = torch.full((12,), int(label), dtype=torch.long)
generated = MODEL.generate(labels, seed=int(seed), temperature=float(temperature))
canvas = Image.new("L", (4 * 128, 3 * 128), color=0)
for index, pixels in enumerate(generated):
image = (
Image.fromarray(
pixels.reshape(8, 8).mul(255).clamp(0, 255).to(torch.uint8).numpy(),
mode="L",
)
.resize((120, 120), Image.Resampling.NEAREST)
)
canvas.paste(image, ((index % 4) * 128 + 4, (index // 4) * 128 + 4))
draw = ImageDraw.Draw(canvas)
for column in range(1, 4):
draw.line((column * 128, 0, column * 128, 384), fill=64, width=1)
for row in range(1, 3):
draw.line((0, row * 128, 512, row * 128), fill=64, width=1)
metadata = {
"digit": int(label),
"seed": int(seed),
"temperature": float(temperature),
"samples": 12,
"mean_pixel_variance": float(generated.var(dim=0).mean()),
}
return canvas, metadata
with gr.Blocks(title="Pocket WGAN-GP") as demo:
gr.Markdown(
"# Pocket WGAN-GP\n"
"Explore a compact adversarial generator trained with Wasserstein distance, "
"gradient penalty, projection conditioning, and explicit collapse checks."
)
with gr.Row():
label = gr.Slider(0, 9, value=7, step=1, label="Digit class")
seed = gr.Slider(0, 100_000, value=2047, step=1, label="Noise seed")
temperature = gr.Slider(
0.25, 1.75, value=1.0, step=0.05, label="Latent temperature"
)
gallery, metadata = generate_gallery(7, 2047, 1.0)
output = gr.Image(value=gallery, label="Twelve adversarial samples")
metrics = gr.JSON(value=metadata, label="Live diversity readout")
button = gr.Button("Generate a new batch", variant="primary")
button.click(
generate_gallery,
inputs=[label, seed, temperature],
outputs=[output, metrics],
)
if __name__ == "__main__":
demo.launch()