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()