from __future__ import annotations from pathlib import Path import gradio as gr import numpy as np import plotly.graph_objects as go from model import RealNVP from safetensors.torch import load_file PROJECT_DIR = Path(__file__).resolve().parent ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "flow-pocket" MODEL = RealNVP() MODEL.load_state_dict(load_file(ARTIFACT_DIR / "realnvp.safetensors")) MODEL.eval() REFERENCE = np.load( ARTIFACT_DIR / "generated_samples.npz" ) def sample_flow(seed: int, temperature: float, samples: int) -> tuple[go.Figure, dict]: generated = MODEL.sample( int(samples), seed=int(seed), temperature=float(temperature) ).numpy() figure = go.Figure() figure.add_trace( go.Scattergl( x=generated[:, 0], y=generated[:, 1], mode="markers", name="RealNVP samples", marker={"size": 4, "opacity": 0.6, "color": "#38bdf8"}, ) ) figure.update_layout( title="Exactly invertible pinwheel generator", xaxis_title="x", yaxis_title="y", template="plotly_dark", yaxis={"scaleanchor": "x", "scaleratio": 1}, ) radius = np.sqrt((generated**2).sum(1)) return figure, { "samples": len(generated), "mean_radius": round(float(radius.mean()), 4), "radius_standard_deviation": round(float(radius.std()), 4), "temperature": float(temperature), } with gr.Blocks(title="Flow Pocket") as demo: gr.Markdown( "# Flow Pocket\n" "Sample an exactly invertible RealNVP and change latent temperature to " "expand or contract the learned pinwheel density." ) with gr.Row(): seed = gr.Number(2043, precision=0, label="Sampling seed") temperature = gr.Slider(0.5, 1.5, 1.0, step=0.05, label="Temperature") samples = gr.Slider(250, 5_000, 2_000, step=250, label="Samples") run = gr.Button("Sample the flow", variant="primary") scatter = gr.Plot() metrics = gr.JSON() run.click(sample_flow, [seed, temperature, samples], [scatter, metrics]) demo.load(sample_flow, [seed, temperature, samples], [scatter, metrics]) if __name__ == "__main__": demo.launch()