Spaces:
Running
Running
Download core/berth_core/pack.py from FineEnvs/PortSimEnv: direct link, hf CLI and curl.
- Browser
- Download file 3.56 kB
-
https://huggingface.co/spaces/FineEnvs/PortSimEnv/resolve/main/core/berth_core/pack.py
- Command line
-
hf download hf://spaces/FineEnvs/PortSimEnv/core/berth_core/pack.py
-
curl -L -o pack.py https://huggingface.co/spaces/FineEnvs/PortSimEnv/resolve/main/core/berth_core/pack.py
3.56 kB
| """Task packs: directories of tasks (tasks.jsonl or tasks.jsonl.gz, with their reference plans) plus a manifest. | |
| The server serves several packs at once; splits come from each task's `split` field. By default that is the dock-v1 | |
| eval and train packs (the hard tasks); the first pack, berth-v1, stays available by setting BERTH_TASKS_DIR. | |
| BERTH_TASKS_DIR=/path/a:/path/b colon-separated pack directories | |
| """ | |
| from __future__ import annotations | |
| import gzip | |
| import hashlib | |
| import json | |
| import os | |
| from pathlib import Path | |
| from .model import Task | |
| TASKS_ROOT = Path(__file__).resolve().parents[2] / "tasks" | |
| DEFAULT_PACKS = ["dock-v1-eval", "dock-v1-train"] | |
| LEGACY_PACK = "berth-v1" | |
| def default_roots() -> list[Path]: | |
| env = os.environ.get("BERTH_TASKS_DIR", "").strip() | |
| if env: | |
| return [Path(p) for p in env.split(":") if p] | |
| roots = [TASKS_ROOT / p for p in DEFAULT_PACKS if _tasks_file(TASKS_ROOT / p)] | |
| return roots or [TASKS_ROOT / LEGACY_PACK] | |
| def _tasks_file(root: Path) -> Path | None: | |
| for name in ("tasks.jsonl", "tasks.jsonl.gz"): | |
| if (root / name).is_file(): | |
| return root / name | |
| return None | |
| def _read(root: Path) -> tuple[list[Task], dict]: | |
| path = _tasks_file(root) | |
| if path is None: | |
| raise FileNotFoundError(f"no tasks.jsonl(.gz) in {root}") | |
| raw = path.read_bytes() | |
| body = gzip.decompress(raw) if path.suffix == ".gz" else raw | |
| manifest = json.loads((root / "manifest.json").read_text()) if (root / "manifest.json").is_file() else {} | |
| digest = hashlib.sha256(body).hexdigest() | |
| if manifest.get("sha256") and manifest["sha256"] != digest: | |
| raise ValueError(f"{path} does not match its manifest.json (sha256 {digest[:12]} != {manifest['sha256'][:12]})") | |
| return [Task.from_dict(json.loads(line)) for line in body.decode().splitlines() if line.strip()], manifest | |
| class TaskPack: | |
| def __init__(self, roots: Path | str | list[Path | str] | None = None): | |
| if roots is None: | |
| roots = default_roots() | |
| if isinstance(roots, (str, Path)): | |
| roots = [roots] | |
| self.roots = [Path(r) for r in roots] | |
| self.root = self.roots[0] | |
| self.tasks: list[Task] = [] | |
| self.manifests: dict[str, dict] = {} | |
| for r in self.roots: | |
| tasks, manifest = _read(r) | |
| self.tasks += tasks | |
| self.manifests[r.name] = manifest | |
| self.manifest = self.manifests[self.root.name] | |
| self._by_id = {} | |
| for t in self.tasks: | |
| if t.task_id in self._by_id: | |
| raise ValueError(f"task id {t.task_id} appears in more than one pack") | |
| self._by_id[t.task_id] = t | |
| self._by_split: dict[str, list[Task]] = {} | |
| for t in self.tasks: | |
| self._by_split.setdefault(t.split, []).append(t) | |
| def splits(self) -> list[str]: | |
| return sorted(self._by_split) | |
| def count(self, split: str) -> int: | |
| return len(self._by_split.get(split, [])) | |
| def at(self, split: str, index: int) -> Task: | |
| return self._by_split[split][index] | |
| def get(self, task_id: str) -> Task: | |
| return self._by_id[task_id] | |
| def public(self, task: Task) -> dict: | |
| """What anyone may see: everything but the reference plans and costs.""" | |
| return task.to_dict(public=True) | |
| _PACK: TaskPack | None = None | |
| def load_pack(root: Path | str | list | None = None) -> TaskPack: | |
| global _PACK | |
| if root is not None: | |
| return TaskPack(root) | |
| if _PACK is None: | |
| _PACK = TaskPack() | |
| return _PACK | |