| |
| """Download the ScienceWorld task-type split used by TCOD / FutureBridge-OPD. |
| |
| Source: https://huggingface.co/datasets/SeanWang0027/scienceworld-tcod-split |
| - train: 2,294 rows / 17 task types |
| - test : 1,308 rows / 13 disjoint task types |
| - columns: task_desc (JSON string), targe (kept for TCOD's loader) |
| |
| The split is portable: every row's jar_path is the empty string, which makes |
| ScienceWorldEnv fall back to the jar bundled with the pip `scienceworld` |
| package -- no machine-local paths, no data regeneration needed. |
| |
| Writes <out>/train.jsonl and <out>/test.jsonl (the layout the released |
| scienceworld YAMLs expect at data/scienceworld/). |
| |
| Usage: |
| python prepare_data.py [--out data/scienceworld] [--smoke] |
| |
| --smoke additionally boots one ScienceWorld episode to verify that the |
| `scienceworld` package and a Java runtime are present. |
| """ |
|
|
| import argparse |
| import json |
| import os |
| import sys |
|
|
| HF_REPO = "SeanWang0027/scienceworld-tcod-split" |
| EXPECTED = {"train": 2294, "test": 1308} |
|
|
|
|
| def main() -> int: |
| parser = argparse.ArgumentParser(description=__doc__) |
| parser.add_argument("--out", default="data/scienceworld") |
| parser.add_argument("--smoke", action="store_true", |
| help="boot one ScienceWorld episode after downloading") |
| args = parser.parse_args() |
|
|
| os.makedirs(args.out, exist_ok=True) |
| for split, expected_rows in EXPECTED.items(): |
| path = os.path.join(args.out, f"{split}.jsonl") |
| if os.path.exists(path): |
| n = sum(1 for _ in open(path)) |
| if n == expected_rows: |
| print(f"[prepare_data] {path} already present ({n} rows), skipping") |
| continue |
| print(f"[prepare_data] {path} has {n} rows, expected {expected_rows}; re-downloading") |
| from datasets import load_dataset |
| ds = load_dataset(HF_REPO, split=split) |
| if len(ds) != expected_rows: |
| print(f"[prepare_data] ERROR: HF split '{split}' has {len(ds)} rows, " |
| f"expected {expected_rows}", file=sys.stderr) |
| return 1 |
| ds.to_json(path, lines=True) |
| print(f"[prepare_data] wrote {path} ({len(ds)} rows)") |
|
|
| |
| with open(os.path.join(args.out, "train.jsonl")) as f: |
| row = json.loads(f.readline()) |
| task_config = json.loads(row["task_desc"]) |
| if task_config.get("jar_path", ""): |
| print(f"[prepare_data] ERROR: row 0 has non-empty jar_path " |
| f"{task_config['jar_path']!r}; this split should be portable", |
| file=sys.stderr) |
| return 1 |
| print(f"[prepare_data] sample task: {task_config['task_name']} " |
| f"var {task_config['var_num']} (jar_path empty -> bundled jar)") |
|
|
| if args.smoke: |
| print("[prepare_data] smoke test: booting one ScienceWorld episode ...") |
| from scienceworld import ScienceWorldEnv |
| env = ScienceWorldEnv("", "", envStepLimit=10) |
| env.load(task_config["task_name"], task_config["var_num"], |
| task_config.get("simplification_str", "easy"), |
| generateGoldPath=False) |
| obs, info = env.reset() |
| obs, reward, done, info = env.step("look around") |
| env.close() |
| print(f"[prepare_data] smoke OK (score={info.get('score')})") |
|
|
| return 0 |
|
|
|
|
| if __name__ == "__main__": |
| sys.exit(main()) |
|
|