ARotting's picture
Publish Prediction-error and learning-progress world models
18603bb verified
Raw
History Blame Contribute Delete
2.26 kB
from __future__ import annotations
from pathlib import Path
import gradio as gr
import numpy as np
import plotly.graph_objects as go
ARTIFACT_DIR = Path(__file__).resolve().parent / "artifacts" / "curiosity-cabinet"
TRAJECTORIES = {
reward: np.load(ARTIFACT_DIR / f"{reward}_trajectory.npz")
for reward in ["prediction_error", "learning_progress"]
}
ROOMS = ["simple", "learnable", "noisy TV"]
def rolling_fraction(actions: np.ndarray, room: int, window: int = 100) -> np.ndarray:
indicator = (actions == room).astype(np.float32)
kernel = np.ones(window, dtype=np.float32) / window
return np.convolve(indicator, kernel, mode="valid")
def compare_curiosity(window: int) -> tuple[go.Figure, dict]:
window = int(window)
figure = go.Figure()
for reward, trajectory in TRAJECTORIES.items():
actions = trajectory["actions"]
for room, room_name in enumerate(ROOMS):
values = rolling_fraction(actions, room, window)
figure.add_trace(
go.Scatter(
x=np.arange(len(values)) + window,
y=values,
mode="lines",
name=f"{reward}: {room_name}",
)
)
figure.update_layout(
title="Where curiosity spends experience",
xaxis_title="Agent step",
yaxis_title="Rolling action fraction",
template="plotly_dark",
)
return figure, {
reward: {
room_name: round(float(np.mean(data["actions"] == room)), 4)
for room, room_name in enumerate(ROOMS)
}
for reward, data in TRAJECTORIES.items()
}
with gr.Blocks(title="Curiosity Cabinet") as demo:
gr.Markdown(
"# Curiosity Cabinet\n"
"Compare surprise-seeking with learning-progress curiosity in a world "
"containing both learnable structure and irreducible noise."
)
window = gr.Slider(25, 250, 100, step=25, label="Rolling window")
run = gr.Button("Replay curious agents", variant="primary")
chart = gr.Plot()
totals = gr.JSON()
run.click(compare_curiosity, window, [chart, totals])
demo.load(compare_curiosity, window, [chart, totals])
if __name__ == "__main__":
demo.launch()