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