ARotting's picture
Publish Offline return-conditioned key-door policy
77a4175 verified
Raw
History Blame Contribute Delete
9.84 kB
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()