ARotting's picture
Publish Offline return-conditioned key-door policy
77a4175 verified
Raw
History Blame Contribute Delete
3.13 kB
from __future__ import annotations
import json
from pathlib import Path
import gradio as gr
import plotly.graph_objects as go
from model import BehaviorCloningPolicy, DecisionTransformer
from safetensors.torch import load_file
from train import rollout_policy
ARTIFACT_DIR = (
Path(__file__).resolve().parent / "artifacts" / "decision-transformer-pocket"
)
DT = DecisionTransformer()
DT.load_state_dict(load_file(ARTIFACT_DIR / "decision_transformer.safetensors"))
DT.eval()
BC = BehaviorCloningPolicy()
BC.load_state_dict(load_file(ARTIFACT_DIR / "behavior_cloning.safetensors"))
BC.eval()
REPORT = json.loads((ARTIFACT_DIR / "evaluation.json").read_text(encoding="utf-8"))
def run_policy(
policy_name: str,
target_return: float,
start: int,
) -> tuple[go.Figure, dict]:
decision_transformer = policy_name == "Decision Transformer"
episode = rollout_policy(
DT if decision_transformer else BC,
decision_transformer=decision_transformer,
target_return=float(target_return),
start=int(start),
)
figure = go.Figure(
go.Scatter(
x=list(range(len(episode["positions"]))),
y=episode["positions"],
mode="lines+markers+text",
text=episode["actions"],
textposition="top center",
)
)
figure.add_hline(y=1, line_dash="dot", annotation_text="key")
figure.add_hline(y=6, line_dash="dot", annotation_text="near reward")
figure.add_hline(y=8, line_dash="dot", annotation_text="treasure")
figure.update_layout(
template="plotly_dark",
title="Offline-RL corridor rollout",
xaxis_title="Step",
yaxis_title="Position",
yaxis_range=[0, 8],
)
key = "decision_transformer" if decision_transformer else "behavior_cloning"
target_key = "treasure_target" if target_return > 0.7 else "near_target"
metrics = {
"terminal": episode["terminal"],
"total_return": episode["total_return"],
"actions": episode["actions"],
"verified_desired_terminal_rate": REPORT["results"][key][target_key][
"desired_terminal_rate"
],
}
return figure, metrics
with gr.Blocks(title="Decision Transformer Pocket") as demo:
gr.Markdown(
"# Decision Transformer Pocket\n"
"Condition an offline policy on desired return: claim the nearby reward or "
"retrieve the key and cross the door for treasure."
)
with gr.Row():
policy = gr.Dropdown(
["Decision Transformer", "Behavior Cloning"],
value="Decision Transformer",
label="Policy",
)
target = gr.Radio([0.4, 1.0], value=1.0, label="Target return")
start = gr.Slider(3, 5, value=4, step=1, label="Start position")
initial = run_policy("Decision Transformer", 1.0, 4)
chart = gr.Plot(value=initial[0])
metrics = gr.JSON(value=initial[1])
button = gr.Button("Run offline policy", variant="primary")
button.click(run_policy, inputs=[policy, target, start], outputs=[chart, metrics])
if __name__ == "__main__":
demo.launch()