AdithyaSK's picture
AdithyaSK HF Staff
Upload folder using huggingface_hub
a1e3907 verified
Raw History Blame Contribute Delete
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