rishavutk commited on
Commit
ea9a79c
·
verified ·
1 Parent(s): 9592f16

Refresh training notebook clone step

Browse files
Files changed (1) hide show
  1. SupplyMind_Training_Run.ipynb +297 -295
SupplyMind_Training_Run.ipynb CHANGED
@@ -1,295 +1,297 @@
1
- {
2
- "cells": [
3
- {
4
- "cell_type": "markdown",
5
- "metadata": {},
6
- "source": [
7
- "# SupplyMind Training Run\n",
8
- "\n",
9
- "This notebook is the compact, judge-runnable training path for SupplyMind: environment smoke test → SFT warm-start → GRPO from SFT → held-out evaluation.\n",
10
- "\n",
11
- "Default settings train the **center** role on the easy task so the notebook can run quickly. Change `ROLE` to `\"warehouse\"` or `TASK_ID` to `\"v2_train_medium\"` / `\"v2_train_hard\"` for a larger run."
12
- ]
13
- },
14
- {
15
- "cell_type": "markdown",
16
- "metadata": {},
17
- "source": [
18
- "## 1. Setup"
19
- ]
20
- },
21
- {
22
- "cell_type": "code",
23
- "execution_count": null,
24
- "metadata": {},
25
- "outputs": [],
26
- "source": [
27
- "!pip -q uninstall -y torchao\n",
28
- "!pip -q install \"torch\" \"transformers>=4.45.0\" \"trl>=0.12.0\" \"peft>=0.13.0\" accelerate datasets bitsandbytes huggingface_hub pydantic pyyaml matplotlib"
29
- ]
30
- },
31
- {
32
- "cell_type": "code",
33
- "execution_count": null,
34
- "metadata": {},
35
- "outputs": [],
36
- "source": [
37
- "import os\n",
38
- "from pathlib import Path\n",
39
- "\n",
40
- "if not Path(\"supplymind\").exists():\n",
41
- " !git clone -q https://huggingface.co/spaces/rishavutk/supplymind supplymind\n",
42
- "\n",
43
- "%cd /content/supplymind\n",
44
- "!pip -q install -e ."
45
- ]
46
- },
47
- {
48
- "cell_type": "code",
49
- "execution_count": null,
50
- "metadata": {},
51
- "outputs": [],
52
- "source": [
53
- "from huggingface_hub import notebook_login, whoami\n",
54
- "\n",
55
- "notebook_login()\n",
56
- "HF_NAMESPACE = whoami()[\"name\"]\n",
57
- "print(\"Using HF namespace:\", HF_NAMESPACE)"
58
- ]
59
- },
60
- {
61
- "cell_type": "markdown",
62
- "metadata": {},
63
- "source": [
64
- "## 2. Controls\n",
65
- "\n",
66
- "Change only these values for quick variants. `ROLE` controls which policy is trained; `TASK_ID` controls easy/medium/hard world generation."
67
- ]
68
- },
69
- {
70
- "cell_type": "code",
71
- "execution_count": null,
72
- "metadata": {},
73
- "outputs": [],
74
- "source": [
75
- "ROLE = \"center\" # \"center\" or \"warehouse\"\n",
76
- "TASK_ID = \"v2_train_easy\" # also: \"v2_train_medium\", \"v2_train_hard\"\n",
77
- "TRAIN_SEEDS = \"101,113,127\"\n",
78
- "EVAL_SEEDS = \"131,149,163\"\n",
79
- "SFT_STEPS = 20\n",
80
- "GRPO_STEPS = 20\n",
81
- "MAX_COMPLETION_LENGTH = 256\n",
82
- "\n",
83
- "SFT_ADAPTER_ID = f\"{HF_NAMESPACE}/supplymind-{ROLE}-qwen-0.5b-sft-notebook\"\n",
84
- "GRPO_ADAPTER_ID = f\"{HF_NAMESPACE}/supplymind-{ROLE}-qwen-0.5b-grpo-notebook\"\n",
85
- "\n",
86
- "print({\n",
87
- " \"role\": ROLE,\n",
88
- " \"task_id\": TASK_ID,\n",
89
- " \"train_seeds\": TRAIN_SEEDS,\n",
90
- " \"eval_seeds\": EVAL_SEEDS,\n",
91
- " \"sft_adapter\": SFT_ADAPTER_ID,\n",
92
- " \"grpo_adapter\": GRPO_ADAPTER_ID,\n",
93
- "})"
94
- ]
95
- },
96
- {
97
- "cell_type": "markdown",
98
- "metadata": {},
99
- "source": [
100
- "## 3. Environment Smoke Test"
101
- ]
102
- },
103
- {
104
- "cell_type": "code",
105
- "execution_count": null,
106
- "metadata": {},
107
- "outputs": [],
108
- "source": [
109
- "import json\n",
110
- "import os\n",
111
- "import sys\n",
112
- "\n",
113
- "sys.path.insert(0, \"/content/supplymind/src\")\n",
114
- "os.environ[\"SUPPLYMIND_REWARD_CONFIG\"] = \"/content/supplymind/configs/supplymind_v2_rewards.yaml\"\n",
115
- "\n",
116
- "from supplymind_env_v2.environment import V2SupplyMindEnv\n",
117
- "from supplymind_env_v2.models import V2JointAction\n",
118
- "\n",
119
- "env = V2SupplyMindEnv(default_task_id=TASK_ID)\n",
120
- "obs = env.reset_internal(TASK_ID, 131)\n",
121
- "print(\"observation keys:\", sorted(obs.model_dump(mode=\"json\").keys()))\n",
122
- "print(\"round:\", obs.round_index, \"warehouses:\", list(obs.warehouses.keys()))\n",
123
- "\n",
124
- "empty_action = {\n",
125
- " \"warehouse_actions\": {},\n",
126
- " \"central_action\": {\n",
127
- " \"central_procurements\": [],\n",
128
- " \"central_liquidations\": [],\n",
129
- " \"central_replenishments\": [],\n",
130
- " \"inventory_transfer_proposals\": [],\n",
131
- " \"offer_matches\": [],\n",
132
- " },\n",
133
- "}\n",
134
- "result = env.step(V2JointAction.model_validate(empty_action))\n",
135
- "print(\"step reward:\", result.reward.step_reward)\n",
136
- "print(\"done:\", result.done)\n",
137
- "print(\"info keys:\", sorted(result.info.keys()))"
138
- ]
139
- },
140
- {
141
- "cell_type": "markdown",
142
- "metadata": {},
143
- "source": [
144
- "## 4. SFT Warm-Start\n",
145
- "\n",
146
- "SFT teaches the model the action JSON shape and a reasonable heuristic policy. For the warehouse role, the notebook enables the conservative SFT flag to reduce invalid or overactive actions."
147
- ]
148
- },
149
- {
150
- "cell_type": "code",
151
- "execution_count": null,
152
- "metadata": {},
153
- "outputs": [],
154
- "source": [
155
- "warehouse_flags = \"--warehouse-conservative-sft --warehouse-signal-limit 2\" if ROLE == \"warehouse\" else \"\"\n",
156
- "\n",
157
- "!python scripts/hf_sft_supplymind_roles.py \\\n",
158
- " --role {ROLE} \\\n",
159
- " --task-id {TASK_ID} \\\n",
160
- " --seeds {TRAIN_SEEDS} \\\n",
161
- " --max-steps {SFT_STEPS} \\\n",
162
- " --hub-model-id {SFT_ADAPTER_ID} \\\n",
163
- " --output-dir outputs/{ROLE}-sft-notebook \\\n",
164
- " {warehouse_flags}"
165
- ]
166
- },
167
- {
168
- "cell_type": "markdown",
169
- "metadata": {},
170
- "source": [
171
- "## 5. GRPO From SFT\n",
172
- "\n",
173
- "The GRPO script routes reward by role: center updates from center reward deltas, warehouse updates from warehouse reward deltas, while global reward is logged for audit."
174
- ]
175
- },
176
- {
177
- "cell_type": "code",
178
- "execution_count": null,
179
- "metadata": {},
180
- "outputs": [],
181
- "source": [
182
- "!python scripts/hf_train_supplymind_roles.py \\\n",
183
- " --role {ROLE} \\\n",
184
- " --task-id {TASK_ID} \\\n",
185
- " --seeds {TRAIN_SEEDS} \\\n",
186
- " --max-steps {GRPO_STEPS} \\\n",
187
- " --max-completion-length {MAX_COMPLETION_LENGTH} \\\n",
188
- " --init-adapter-id {SFT_ADAPTER_ID} \\\n",
189
- " --hub-model-id {GRPO_ADAPTER_ID} \\\n",
190
- " --output-dir outputs/{ROLE}-grpo-notebook"
191
- ]
192
- },
193
- {
194
- "cell_type": "markdown",
195
- "metadata": {},
196
- "source": [
197
- "## 6. Held-Out Evaluation\n",
198
- "\n",
199
- "Evaluate base Qwen, SFT, and GRPO on held-out seeds. The role score is the training target; global score is the environment-level audit metric."
200
- ]
201
- },
202
- {
203
- "cell_type": "code",
204
- "execution_count": null,
205
- "metadata": {},
206
- "outputs": [],
207
- "source": [
208
- "!python scripts/hf_eval_supplymind_adapters.py \\\n",
209
- " --role {ROLE} \\\n",
210
- " --task-id {TASK_ID} \\\n",
211
- " --seeds {EVAL_SEEDS} \\\n",
212
- " --sft-adapter-id {SFT_ADAPTER_ID} \\\n",
213
- " --grpo-adapter-id {GRPO_ADAPTER_ID} \\\n",
214
- " --max-new-tokens {MAX_COMPLETION_LENGTH} | tee outputs/{ROLE}-eval-notebook.log"
215
- ]
216
- },
217
- {
218
- "cell_type": "code",
219
- "execution_count": null,
220
- "metadata": {},
221
- "outputs": [],
222
- "source": [
223
- "import json\n",
224
- "import re\n",
225
- "from pathlib import Path\n",
226
- "\n",
227
- "import matplotlib.pyplot as plt\n",
228
- "import pandas as pd\n",
229
- "\n",
230
- "log_path = Path(f\"outputs/{ROLE}-eval-notebook.log\")\n",
231
- "rows = []\n",
232
- "for line in log_path.read_text(encoding=\"utf-8\", errors=\"ignore\").splitlines():\n",
233
- " line = line.strip()\n",
234
- " if not line.startswith(\"{\"):\n",
235
- " continue\n",
236
- " try:\n",
237
- " payload = json.loads(line)\n",
238
- " except json.JSONDecodeError:\n",
239
- " continue\n",
240
- " if payload.get(\"message\") == \"eval_done\":\n",
241
- " evaluations = {key: payload[key] for key in (\"base\", \"sft\", \"grpo\") if key in payload}\n",
242
- " for label, item in evaluations.items():\n",
243
- " role_score_key = \"mean_center_role_score\" if ROLE == \"center\" else \"mean_warehouse_role_score\"\n",
244
- " rows.append({\n",
245
- " \"policy\": label,\n",
246
- " \"global_score\": item.get(\"mean_global_score\"),\n",
247
- " \"role_score\": item.get(role_score_key),\n",
248
- " \"raw_reward\": item.get(\"mean_raw_reward\"),\n",
249
- " \"invalid_payloads\": item.get(\"invalid_payloads\"),\n",
250
- " \"invalid_actions\": item.get(\"invalid_actions\"),\n",
251
- " })\n",
252
- "\n",
253
- "df = pd.DataFrame(rows)\n",
254
- "display(df)\n",
255
- "\n",
256
- "if not df.empty:\n",
257
- " fig, axes = plt.subplots(1, 2, figsize=(10, 4))\n",
258
- " df.plot.bar(x=\"policy\", y=\"role_score\", ax=axes[0], legend=False, color=\"#2563eb\")\n",
259
- " axes[0].set_title(f\"{ROLE} role score\")\n",
260
- " axes[0].set_ylim(0, 1)\n",
261
- " df.plot.bar(x=\"policy\", y=[\"invalid_payloads\", \"invalid_actions\"], ax=axes[1], color=[\"#dc2626\", \"#f59e0b\"])\n",
262
- " axes[1].set_title(\"Invalid outputs\")\n",
263
- " plt.tight_layout()\n",
264
- " plt.show()"
265
- ]
266
- },
267
- {
268
- "cell_type": "markdown",
269
- "metadata": {},
270
- "source": [
271
- "## Rerun For The Other Role\n",
272
- "\n",
273
- "To train warehouses instead of center, change `ROLE = \"warehouse\"` in the controls cell and rerun sections 4-6. To try larger worlds, change `TASK_ID` to `v2_train_medium` or `v2_train_hard`."
274
- ]
275
- }
276
- ],
277
- "metadata": {
278
- "accelerator": "GPU",
279
- "colab": {
280
- "gpuType": "T4",
281
- "provenance": []
282
- },
283
- "kernelspec": {
284
- "display_name": "Python 3",
285
- "language": "python",
286
- "name": "python3"
287
- },
288
- "language_info": {
289
- "name": "python",
290
- "version": "3.11"
291
- }
292
- },
293
- "nbformat": 4,
294
- "nbformat_minor": 5
295
- }
 
 
 
1
+ {
2
+ "cells": [
3
+ {
4
+ "cell_type": "markdown",
5
+ "metadata": {},
6
+ "source": [
7
+ "# SupplyMind Training Run\n",
8
+ "\n",
9
+ "This notebook is the compact, judge-runnable training path for SupplyMind: environment smoke test → SFT warm-start → GRPO from SFT → held-out evaluation.\n",
10
+ "\n",
11
+ "Default settings train the **center** role on the easy task so the notebook can run quickly. Change `ROLE` to `\"warehouse\"` or `TASK_ID` to `\"v2_train_medium\"` / `\"v2_train_hard\"` for a larger run."
12
+ ]
13
+ },
14
+ {
15
+ "cell_type": "markdown",
16
+ "metadata": {},
17
+ "source": [
18
+ "## 1. Setup"
19
+ ]
20
+ },
21
+ {
22
+ "cell_type": "code",
23
+ "execution_count": null,
24
+ "metadata": {},
25
+ "outputs": [],
26
+ "source": [
27
+ "!pip -q uninstall -y torchao\n",
28
+ "!pip -q install \"torch\" \"transformers>=4.45.0\" \"trl>=0.12.0\" \"peft>=0.13.0\" accelerate datasets bitsandbytes huggingface_hub pydantic pyyaml matplotlib"
29
+ ]
30
+ },
31
+ {
32
+ "cell_type": "code",
33
+ "execution_count": null,
34
+ "metadata": {},
35
+ "outputs": [],
36
+ "source": [
37
+ "import os\n",
38
+ "from pathlib import Path\n",
39
+ "\n",
40
+ "if not Path(\"supplymind\").exists():\n",
41
+ " !git clone -q https://huggingface.co/spaces/rishavutk/supplymind supplymind\n",
42
+ "else:\n",
43
+ " !git -C supplymind pull -q\n",
44
+ "\n",
45
+ "%cd /content/supplymind\n",
46
+ "!pip -q install -e .\n"
47
+ ]
48
+ },
49
+ {
50
+ "cell_type": "code",
51
+ "execution_count": null,
52
+ "metadata": {},
53
+ "outputs": [],
54
+ "source": [
55
+ "from huggingface_hub import notebook_login, whoami\n",
56
+ "\n",
57
+ "notebook_login()\n",
58
+ "HF_NAMESPACE = whoami()[\"name\"]\n",
59
+ "print(\"Using HF namespace:\", HF_NAMESPACE)"
60
+ ]
61
+ },
62
+ {
63
+ "cell_type": "markdown",
64
+ "metadata": {},
65
+ "source": [
66
+ "## 2. Controls\n",
67
+ "\n",
68
+ "Change only these values for quick variants. `ROLE` controls which policy is trained; `TASK_ID` controls easy/medium/hard world generation."
69
+ ]
70
+ },
71
+ {
72
+ "cell_type": "code",
73
+ "execution_count": null,
74
+ "metadata": {},
75
+ "outputs": [],
76
+ "source": [
77
+ "ROLE = \"center\" # \"center\" or \"warehouse\"\n",
78
+ "TASK_ID = \"v2_train_easy\" # also: \"v2_train_medium\", \"v2_train_hard\"\n",
79
+ "TRAIN_SEEDS = \"101,113,127\"\n",
80
+ "EVAL_SEEDS = \"131,149,163\"\n",
81
+ "SFT_STEPS = 20\n",
82
+ "GRPO_STEPS = 20\n",
83
+ "MAX_COMPLETION_LENGTH = 256\n",
84
+ "\n",
85
+ "SFT_ADAPTER_ID = f\"{HF_NAMESPACE}/supplymind-{ROLE}-qwen-0.5b-sft-notebook\"\n",
86
+ "GRPO_ADAPTER_ID = f\"{HF_NAMESPACE}/supplymind-{ROLE}-qwen-0.5b-grpo-notebook\"\n",
87
+ "\n",
88
+ "print({\n",
89
+ " \"role\": ROLE,\n",
90
+ " \"task_id\": TASK_ID,\n",
91
+ " \"train_seeds\": TRAIN_SEEDS,\n",
92
+ " \"eval_seeds\": EVAL_SEEDS,\n",
93
+ " \"sft_adapter\": SFT_ADAPTER_ID,\n",
94
+ " \"grpo_adapter\": GRPO_ADAPTER_ID,\n",
95
+ "})"
96
+ ]
97
+ },
98
+ {
99
+ "cell_type": "markdown",
100
+ "metadata": {},
101
+ "source": [
102
+ "## 3. Environment Smoke Test"
103
+ ]
104
+ },
105
+ {
106
+ "cell_type": "code",
107
+ "execution_count": null,
108
+ "metadata": {},
109
+ "outputs": [],
110
+ "source": [
111
+ "import json\n",
112
+ "import os\n",
113
+ "import sys\n",
114
+ "\n",
115
+ "sys.path.insert(0, \"/content/supplymind/src\")\n",
116
+ "os.environ[\"SUPPLYMIND_REWARD_CONFIG\"] = \"/content/supplymind/configs/supplymind_v2_rewards.yaml\"\n",
117
+ "\n",
118
+ "from supplymind_env_v2.environment import V2SupplyMindEnv\n",
119
+ "from supplymind_env_v2.models import V2JointAction\n",
120
+ "\n",
121
+ "env = V2SupplyMindEnv(default_task_id=TASK_ID)\n",
122
+ "obs = env.reset_internal(TASK_ID, 131)\n",
123
+ "print(\"observation keys:\", sorted(obs.model_dump(mode=\"json\").keys()))\n",
124
+ "print(\"round:\", obs.round_index, \"warehouses:\", list(obs.warehouses.keys()))\n",
125
+ "\n",
126
+ "empty_action = {\n",
127
+ " \"warehouse_actions\": {},\n",
128
+ " \"central_action\": {\n",
129
+ " \"central_procurements\": [],\n",
130
+ " \"central_liquidations\": [],\n",
131
+ " \"central_replenishments\": [],\n",
132
+ " \"inventory_transfer_proposals\": [],\n",
133
+ " \"offer_matches\": [],\n",
134
+ " },\n",
135
+ "}\n",
136
+ "result = env.step(V2JointAction.model_validate(empty_action))\n",
137
+ "print(\"step reward:\", result.reward.step_reward)\n",
138
+ "print(\"done:\", result.done)\n",
139
+ "print(\"info keys:\", sorted(result.info.keys()))"
140
+ ]
141
+ },
142
+ {
143
+ "cell_type": "markdown",
144
+ "metadata": {},
145
+ "source": [
146
+ "## 4. SFT Warm-Start\n",
147
+ "\n",
148
+ "SFT teaches the model the action JSON shape and a reasonable heuristic policy. For the warehouse role, the notebook enables the conservative SFT flag to reduce invalid or overactive actions."
149
+ ]
150
+ },
151
+ {
152
+ "cell_type": "code",
153
+ "execution_count": null,
154
+ "metadata": {},
155
+ "outputs": [],
156
+ "source": [
157
+ "warehouse_flags = \"--warehouse-conservative-sft --warehouse-signal-limit 2\" if ROLE == \"warehouse\" else \"\"\n",
158
+ "\n",
159
+ "!python scripts/hf_sft_supplymind_roles.py \\\n",
160
+ " --role {ROLE} \\\n",
161
+ " --task-id {TASK_ID} \\\n",
162
+ " --seeds {TRAIN_SEEDS} \\\n",
163
+ " --max-steps {SFT_STEPS} \\\n",
164
+ " --hub-model-id {SFT_ADAPTER_ID} \\\n",
165
+ " --output-dir outputs/{ROLE}-sft-notebook \\\n",
166
+ " {warehouse_flags}"
167
+ ]
168
+ },
169
+ {
170
+ "cell_type": "markdown",
171
+ "metadata": {},
172
+ "source": [
173
+ "## 5. GRPO From SFT\n",
174
+ "\n",
175
+ "The GRPO script routes reward by role: center updates from center reward deltas, warehouse updates from warehouse reward deltas, while global reward is logged for audit."
176
+ ]
177
+ },
178
+ {
179
+ "cell_type": "code",
180
+ "execution_count": null,
181
+ "metadata": {},
182
+ "outputs": [],
183
+ "source": [
184
+ "!python scripts/hf_train_supplymind_roles.py \\\n",
185
+ " --role {ROLE} \\\n",
186
+ " --task-id {TASK_ID} \\\n",
187
+ " --seeds {TRAIN_SEEDS} \\\n",
188
+ " --max-steps {GRPO_STEPS} \\\n",
189
+ " --max-completion-length {MAX_COMPLETION_LENGTH} \\\n",
190
+ " --init-adapter-id {SFT_ADAPTER_ID} \\\n",
191
+ " --hub-model-id {GRPO_ADAPTER_ID} \\\n",
192
+ " --output-dir outputs/{ROLE}-grpo-notebook"
193
+ ]
194
+ },
195
+ {
196
+ "cell_type": "markdown",
197
+ "metadata": {},
198
+ "source": [
199
+ "## 6. Held-Out Evaluation\n",
200
+ "\n",
201
+ "Evaluate base Qwen, SFT, and GRPO on held-out seeds. The role score is the training target; global score is the environment-level audit metric."
202
+ ]
203
+ },
204
+ {
205
+ "cell_type": "code",
206
+ "execution_count": null,
207
+ "metadata": {},
208
+ "outputs": [],
209
+ "source": [
210
+ "!python scripts/hf_eval_supplymind_adapters.py \\\n",
211
+ " --role {ROLE} \\\n",
212
+ " --task-id {TASK_ID} \\\n",
213
+ " --seeds {EVAL_SEEDS} \\\n",
214
+ " --sft-adapter-id {SFT_ADAPTER_ID} \\\n",
215
+ " --grpo-adapter-id {GRPO_ADAPTER_ID} \\\n",
216
+ " --max-new-tokens {MAX_COMPLETION_LENGTH} | tee outputs/{ROLE}-eval-notebook.log"
217
+ ]
218
+ },
219
+ {
220
+ "cell_type": "code",
221
+ "execution_count": null,
222
+ "metadata": {},
223
+ "outputs": [],
224
+ "source": [
225
+ "import json\n",
226
+ "import re\n",
227
+ "from pathlib import Path\n",
228
+ "\n",
229
+ "import matplotlib.pyplot as plt\n",
230
+ "import pandas as pd\n",
231
+ "\n",
232
+ "log_path = Path(f\"outputs/{ROLE}-eval-notebook.log\")\n",
233
+ "rows = []\n",
234
+ "for line in log_path.read_text(encoding=\"utf-8\", errors=\"ignore\").splitlines():\n",
235
+ " line = line.strip()\n",
236
+ " if not line.startswith(\"{\"):\n",
237
+ " continue\n",
238
+ " try:\n",
239
+ " payload = json.loads(line)\n",
240
+ " except json.JSONDecodeError:\n",
241
+ " continue\n",
242
+ " if payload.get(\"message\") == \"eval_done\":\n",
243
+ " evaluations = {key: payload[key] for key in (\"base\", \"sft\", \"grpo\") if key in payload}\n",
244
+ " for label, item in evaluations.items():\n",
245
+ " role_score_key = \"mean_center_role_score\" if ROLE == \"center\" else \"mean_warehouse_role_score\"\n",
246
+ " rows.append({\n",
247
+ " \"policy\": label,\n",
248
+ " \"global_score\": item.get(\"mean_global_score\"),\n",
249
+ " \"role_score\": item.get(role_score_key),\n",
250
+ " \"raw_reward\": item.get(\"mean_raw_reward\"),\n",
251
+ " \"invalid_payloads\": item.get(\"invalid_payloads\"),\n",
252
+ " \"invalid_actions\": item.get(\"invalid_actions\"),\n",
253
+ " })\n",
254
+ "\n",
255
+ "df = pd.DataFrame(rows)\n",
256
+ "display(df)\n",
257
+ "\n",
258
+ "if not df.empty:\n",
259
+ " fig, axes = plt.subplots(1, 2, figsize=(10, 4))\n",
260
+ " df.plot.bar(x=\"policy\", y=\"role_score\", ax=axes[0], legend=False, color=\"#2563eb\")\n",
261
+ " axes[0].set_title(f\"{ROLE} role score\")\n",
262
+ " axes[0].set_ylim(0, 1)\n",
263
+ " df.plot.bar(x=\"policy\", y=[\"invalid_payloads\", \"invalid_actions\"], ax=axes[1], color=[\"#dc2626\", \"#f59e0b\"])\n",
264
+ " axes[1].set_title(\"Invalid outputs\")\n",
265
+ " plt.tight_layout()\n",
266
+ " plt.show()"
267
+ ]
268
+ },
269
+ {
270
+ "cell_type": "markdown",
271
+ "metadata": {},
272
+ "source": [
273
+ "## Rerun For The Other Role\n",
274
+ "\n",
275
+ "To train warehouses instead of center, change `ROLE = \"warehouse\"` in the controls cell and rerun sections 4-6. To try larger worlds, change `TASK_ID` to `v2_train_medium` or `v2_train_hard`."
276
+ ]
277
+ }
278
+ ],
279
+ "metadata": {
280
+ "accelerator": "GPU",
281
+ "colab": {
282
+ "gpuType": "T4",
283
+ "provenance": []
284
+ },
285
+ "kernelspec": {
286
+ "display_name": "Python 3",
287
+ "language": "python",
288
+ "name": "python3"
289
+ },
290
+ "language_info": {
291
+ "name": "python",
292
+ "version": "3.11"
293
+ }
294
+ },
295
+ "nbformat": 4,
296
+ "nbformat_minor": 5
297
+ }