| from __future__ import annotations |
|
|
| from collections import deque |
|
|
| import numpy as np |
| import torch |
| from torch import nn |
|
|
| ROOMS = ["simple", "learnable", "noisy_tv"] |
|
|
|
|
| class WorldModel(nn.Module): |
| def __init__(self) -> None: |
| super().__init__() |
| self.network = nn.Sequential( |
| nn.Linear(1, 16), |
| nn.Tanh(), |
| nn.Linear(16, 1), |
| ) |
|
|
| def forward(self, context: torch.Tensor) -> torch.Tensor: |
| return self.network(context) |
|
|
|
|
| def outcome(room: int, context: float, rng: np.random.Generator) -> float: |
| if room == 0: |
| return 0.25 |
| if room == 1: |
| return float(np.sin(3.0 * context) + 0.45 * context) |
| return float(rng.normal()) |
|
|
|
|
| @torch.inference_mode() |
| def learnable_mse(model: WorldModel) -> float: |
| grid = torch.linspace(-1, 1, 256).unsqueeze(1) |
| truth = torch.sin(3 * grid) + 0.45 * grid |
| return float(torch.mean((model(grid) - truth) ** 2)) |
|
|
|
|
| def intrinsic_scores( |
| histories: list[deque], reward: str |
| ) -> np.ndarray: |
| scores = np.zeros(len(histories), dtype=np.float64) |
| for room, history in enumerate(histories): |
| values = np.asarray(history, dtype=np.float64) |
| if len(values) < 40: |
| scores[room] = 0.0 |
| elif reward == "prediction_error": |
| scores[room] = values[-20:].mean() |
| else: |
| scores[room] = max( |
| 0.0, values[-40:-20].mean() - values[-20:].mean() |
| ) |
| return scores |
|
|
|
|
| def run_agent( |
| reward: str, |
| seed: int, |
| steps: int = 1_200, |
| epsilon: float = 0.15, |
| ) -> dict: |
| if reward not in {"prediction_error", "learning_progress"}: |
| raise ValueError(reward) |
| torch.manual_seed(seed) |
| rng = np.random.default_rng(seed) |
| models = [WorldModel() for _ in ROOMS] |
| optimizers = [ |
| torch.optim.SGD(model.parameters(), lr=0.035) for model in models |
| ] |
| histories = [deque(maxlen=40) for _ in ROOMS] |
| actions = [] |
| losses = [] |
| for step in range(steps): |
| scores = intrinsic_scores(histories, reward) |
| if step < 120 or rng.random() < epsilon: |
| room = int(rng.integers(0, len(ROOMS))) |
| else: |
| room = int(np.argmax(scores + rng.normal(scale=1e-8, size=3))) |
| context = float(rng.uniform(-1, 1)) |
| target = outcome(room, context, rng) |
| prediction = models[room](torch.tensor([[context]], dtype=torch.float32)) |
| loss = (prediction.squeeze() - target) ** 2 |
| optimizers[room].zero_grad() |
| loss.backward() |
| torch.nn.utils.clip_grad_norm_(models[room].parameters(), 2.0) |
| optimizers[room].step() |
| value = float(loss.detach()) |
| histories[room].append(value) |
| actions.append(room) |
| losses.append(value) |
| actions_array = np.asarray(actions) |
|
|
| def fractions(start: int, end: int) -> dict: |
| window = actions_array[start:end] |
| return { |
| room: float(np.mean(window == index)) |
| for index, room in enumerate(ROOMS) |
| } |
|
|
| return { |
| "reward": reward, |
| "seed": seed, |
| "steps": steps, |
| "overall_action_fraction": fractions(0, steps), |
| "middle_action_fraction": fractions(steps // 4, 3 * steps // 4), |
| "final_action_fraction": fractions(3 * steps // 4, steps), |
| "learnable_world_model_mse": learnable_mse(models[1]), |
| "actions": actions, |
| "losses": losses, |
| "learnable_state_dict": models[1].state_dict(), |
| } |
|
|