| from __future__ import annotations |
|
|
| import json |
| import random |
| from pathlib import Path |
|
|
| import numpy as np |
| import pandas as pd |
| import torch |
| import trackio |
| from environment import ACTIONS, KeyDoorCorridor, expert_action |
| from model import BehaviorCloningPolicy, DecisionTransformer, parameter_count |
| from safetensors.torch import save_file |
| from torch.nn import functional as F |
|
|
| PROJECT_DIR = Path(__file__).resolve().parent |
| ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "decision-transformer-pocket" |
| DATA_DIR = PROJECT_DIR / "data" |
| SEED = 2221 |
| MAXIMUM_LENGTH = 20 |
|
|
|
|
| def generate_episode(rng: np.random.Generator) -> dict: |
| start = int(rng.integers(3, 6)) |
| kind = rng.choice(["treasure", "near_reward", "random"], p=[0.4, 0.4, 0.2]) |
| environment = KeyDoorCorridor(start, MAXIMUM_LENGTH) |
| positions, keys, actions, rewards = [], [], [], [] |
| while not environment.done: |
| positions.append(environment.position) |
| keys.append(int(environment.has_key)) |
| if kind == "random" or rng.random() < 0.12: |
| action = int(rng.integers(0, 3)) |
| else: |
| action = expert_action(environment, kind) |
| actions.append(action) |
| rewards.append(environment.step(action)) |
| returns_to_go = np.cumsum(rewards[::-1])[::-1].astype(np.float32) |
| previous_actions = [3] + actions[:-1] |
| return { |
| "start": start, |
| "kind": str(kind), |
| "terminal": environment.terminal, |
| "positions": positions, |
| "keys": keys, |
| "actions": actions, |
| "previous_actions": previous_actions, |
| "rewards": rewards, |
| "returns_to_go": returns_to_go.tolist(), |
| "total_return": float(sum(rewards)), |
| } |
|
|
|
|
| def pad_episode(episode: dict) -> tuple[torch.Tensor, ...]: |
| length = len(episode["actions"]) |
| padding = MAXIMUM_LENGTH - length |
| return ( |
| torch.tensor(episode["positions"] + [0] * padding), |
| torch.tensor(episode["keys"] + [0] * padding), |
| torch.tensor(episode["returns_to_go"] + [0.0] * padding), |
| torch.tensor(episode["previous_actions"] + [3] * padding), |
| torch.tensor(episode["actions"] + [-100] * padding), |
| torch.tensor([True] * length + [False] * padding), |
| ) |
|
|
|
|
| def rollout_policy( |
| model: torch.nn.Module, |
| *, |
| decision_transformer: bool, |
| target_return: float, |
| start: int, |
| ) -> dict: |
| environment = KeyDoorCorridor(start, MAXIMUM_LENGTH) |
| positions, keys, previous_actions, returns = [], [], [], [] |
| actions = [] |
| rewards = [] |
| remaining_return = target_return |
| previous_action = 3 |
| with torch.inference_mode(): |
| while not environment.done: |
| positions.append(environment.position) |
| keys.append(int(environment.has_key)) |
| previous_actions.append(previous_action) |
| returns.append(remaining_return) |
| if decision_transformer: |
| length = len(positions) |
| logits = model( |
| torch.tensor([positions]), |
| torch.tensor([keys]), |
| torch.tensor([returns]), |
| torch.tensor([previous_actions]), |
| torch.ones(1, length, dtype=torch.bool), |
| )[0, -1] |
| else: |
| logits = model( |
| torch.tensor([environment.position]), |
| torch.tensor([int(environment.has_key)]), |
| )[0] |
| action = int(logits.argmax()) |
| reward = environment.step(action) |
| actions.append(action) |
| rewards.append(reward) |
| remaining_return -= reward |
| previous_action = action |
| return { |
| "start": start, |
| "target_return": target_return, |
| "terminal": environment.terminal, |
| "total_return": float(sum(rewards)), |
| "positions": positions, |
| "actions": [ACTIONS[action] for action in actions], |
| "rewards": rewards, |
| } |
|
|
|
|
| def evaluate(model: torch.nn.Module, *, decision_transformer: bool) -> dict: |
| report = {} |
| for name, target, expected in [ |
| ("near_target", 0.4, "near_reward"), |
| ("treasure_target", 1.0, "treasure"), |
| ]: |
| episodes = [ |
| rollout_policy( |
| model, |
| decision_transformer=decision_transformer, |
| target_return=target, |
| start=start, |
| ) |
| for start in [3, 4, 5] |
| for _ in range(100) |
| ] |
| report[name] = { |
| "desired_terminal": expected, |
| "desired_terminal_rate": sum( |
| episode["terminal"] == expected for episode in episodes |
| ) |
| / len(episodes), |
| "mean_return": float( |
| np.mean([episode["total_return"] for episode in episodes]) |
| ), |
| "mean_steps": float(np.mean([len(episode["actions"]) for episode in episodes])), |
| "episodes": len(episodes), |
| } |
| return report |
|
|
|
|
| def main() -> None: |
| random.seed(SEED) |
| np.random.seed(SEED) |
| torch.manual_seed(SEED) |
| torch.set_num_threads(1) |
| rng = np.random.default_rng(SEED) |
| episodes = [generate_episode(rng) for _ in range(4_000)] |
| padded = [pad_episode(episode) for episode in episodes] |
| tensors = [torch.stack(items) for items in zip(*padded, strict=True)] |
| positions, keys, returns, previous_actions, action_targets, valid = tensors |
| decision_transformer = DecisionTransformer() |
| behavior_cloning = BehaviorCloningPolicy() |
| dt_optimizer = torch.optim.AdamW( |
| decision_transformer.parameters(), lr=2e-3, weight_decay=1e-4 |
| ) |
| bc_optimizer = torch.optim.AdamW( |
| behavior_cloning.parameters(), lr=2e-3, weight_decay=1e-4 |
| ) |
| trackio.init( |
| project="decision-transformer-pocket", |
| name="key-door-return-conditioning-v1", |
| config={ |
| "decision_transformer_parameters": parameter_count(decision_transformer), |
| "behavior_cloning_parameters": parameter_count(behavior_cloning), |
| "offline_episodes": len(episodes), |
| "epochs": 80, |
| }, |
| ) |
| for epoch in range(1, 81): |
| order = torch.randperm(len(episodes)) |
| decision_transformer.train() |
| behavior_cloning.train() |
| for start in range(0, len(episodes), 128): |
| indexes = order[start : start + 128] |
| dt_logits = decision_transformer( |
| positions[indexes], |
| keys[indexes], |
| returns[indexes], |
| previous_actions[indexes], |
| valid[indexes], |
| ) |
| dt_loss = F.cross_entropy( |
| dt_logits.flatten(0, 1), |
| action_targets[indexes].flatten(), |
| ignore_index=-100, |
| ) |
| dt_optimizer.zero_grad(set_to_none=True) |
| dt_loss.backward() |
| dt_optimizer.step() |
| mask = valid[indexes] |
| bc_logits = behavior_cloning( |
| positions[indexes][mask], keys[indexes][mask] |
| ) |
| bc_loss = F.cross_entropy(bc_logits, action_targets[indexes][mask]) |
| bc_optimizer.zero_grad(set_to_none=True) |
| bc_loss.backward() |
| bc_optimizer.step() |
| if epoch == 1 or epoch % 10 == 0: |
| trackio.log( |
| { |
| "epoch": epoch, |
| "decision_transformer_loss": float(dt_loss.detach()), |
| "behavior_cloning_loss": float(bc_loss.detach()), |
| } |
| ) |
| decision_transformer.eval() |
| behavior_cloning.eval() |
| results = { |
| "decision_transformer": { |
| "parameters": parameter_count(decision_transformer), |
| **evaluate(decision_transformer, decision_transformer=True), |
| }, |
| "behavior_cloning": { |
| "parameters": parameter_count(behavior_cloning), |
| **evaluate(behavior_cloning, decision_transformer=False), |
| }, |
| } |
| report = { |
| "experiment": "Offline return-conditioned key-door control", |
| "offline_episodes": len(episodes), |
| "results": results, |
| } |
| ARTIFACT_DIR.mkdir(parents=True, exist_ok=True) |
| DATA_DIR.mkdir(parents=True, exist_ok=True) |
| save_file( |
| decision_transformer.state_dict(), |
| ARTIFACT_DIR / "decision_transformer.safetensors", |
| ) |
| save_file( |
| behavior_cloning.state_dict(), |
| ARTIFACT_DIR / "behavior_cloning.safetensors", |
| ) |
| (ARTIFACT_DIR / "evaluation.json").write_text( |
| json.dumps(report, indent=2), encoding="utf-8" |
| ) |
| pd.DataFrame( |
| [ |
| { |
| **episode, |
| "positions": json.dumps(episode["positions"]), |
| "keys": json.dumps(episode["keys"]), |
| "actions": json.dumps(episode["actions"]), |
| "previous_actions": json.dumps(episode["previous_actions"]), |
| "rewards": json.dumps(episode["rewards"]), |
| "returns_to_go": json.dumps(episode["returns_to_go"]), |
| } |
| for episode in episodes |
| ] |
| ).to_parquet(DATA_DIR / "offline_trajectories.parquet", index=False) |
| trackio.log( |
| { |
| "dt_near_rate": results["decision_transformer"]["near_target"][ |
| "desired_terminal_rate" |
| ], |
| "dt_treasure_rate": results["decision_transformer"]["treasure_target"][ |
| "desired_terminal_rate" |
| ], |
| "bc_near_rate": results["behavior_cloning"]["near_target"][ |
| "desired_terminal_rate" |
| ], |
| "bc_treasure_rate": results["behavior_cloning"]["treasure_target"][ |
| "desired_terminal_rate" |
| ], |
| } |
| ) |
| trackio.finish() |
| print(json.dumps(report, indent=2)) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|