Download scripts/split_task_stream.py from RL-MIND/Harness-Continual-Learning: direct link, hf CLI and curl.
- Browser
- Download file 6.14 kB
-
https://huggingface.co/RL-MIND/Harness-Continual-Learning/resolve/main/scripts/split_task_stream.py
- Command line
-
hf download hf://RL-MIND/Harness-Continual-Learning/scripts/split_task_stream.py
-
curl -L -o split_task_stream.py https://huggingface.co/RL-MIND/Harness-Continual-Learning/resolve/main/scripts/split_task_stream.py
6.14 kB
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import random | |
| from pathlib import Path | |
| from typing import Any | |
| def read_records(path: Path) -> tuple[list[dict[str, Any]], str]: | |
| if path.suffix == ".jsonl": | |
| rows = [ | |
| json.loads(line) | |
| for line in path.read_text(encoding="utf-8").splitlines() | |
| if line.strip() | |
| ] | |
| return [dict(row) for row in rows if isinstance(row, dict)], "jsonl" | |
| raw = json.loads(path.read_text(encoding="utf-8")) | |
| if isinstance(raw, list): | |
| return [dict(row) for row in raw if isinstance(row, dict)], "list" | |
| if isinstance(raw, dict): | |
| for key in ("data", "examples", "instances"): | |
| value = raw.get(key) | |
| if isinstance(value, list): | |
| return [dict(row) for row in value if isinstance(row, dict)], key | |
| raise ValueError(f"Unsupported dataset shape: {path}") | |
| def write_records(path: Path, records: list[dict[str, Any]], shape: str) -> None: | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| if shape == "jsonl": | |
| path.write_text( | |
| "".join(json.dumps(record, ensure_ascii=False) + "\n" for record in records), | |
| encoding="utf-8", | |
| ) | |
| return | |
| payload: Any = records if shape == "list" else {shape: records} | |
| path.write_text(json.dumps(payload, ensure_ascii=False, indent=2) + "\n", encoding="utf-8") | |
| def split_records( | |
| records: list[dict[str, Any]], | |
| *, | |
| val_ratio: float, | |
| min_val: int, | |
| max_val: int | None, | |
| seed: int, | |
| ) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]: | |
| if not records: | |
| return [], [] | |
| shuffled = list(records) | |
| random.Random(seed).shuffle(shuffled) | |
| val_size = max(min_val, int(round(len(shuffled) * val_ratio))) | |
| if max_val is not None: | |
| val_size = min(val_size, max_val) | |
| val_size = max(1, min(val_size, len(shuffled) - 1)) if len(shuffled) > 1 else 1 | |
| return shuffled[val_size:], shuffled[:val_size] | |
| def build_parser() -> argparse.ArgumentParser: | |
| parser = argparse.ArgumentParser( | |
| description="Split each task's train file into deterministic train and validation files." | |
| ) | |
| parser.add_argument("--config", type=Path, required=True, help="Task-stream config to read.") | |
| parser.add_argument("--output-config", type=Path, help="Updated task-stream config to write.") | |
| parser.add_argument("--output-dir", type=Path, help="Common output directory for split files.") | |
| parser.add_argument("--val-ratio", type=float, default=0.2) | |
| parser.add_argument("--min-val", type=int, default=1) | |
| parser.add_argument("--max-val", type=int) | |
| parser.add_argument("--seed", type=int, default=42) | |
| parser.add_argument("--overwrite", action="store_true") | |
| parser.add_argument("--dry-run", action="store_true") | |
| return parser | |
| def main() -> int: | |
| args = build_parser().parse_args() | |
| if not 0 < args.val_ratio < 1: | |
| raise ValueError("--val-ratio must be between 0 and 1.") | |
| if args.min_val < 0 or (args.max_val is not None and args.max_val < 1): | |
| raise ValueError("Validation limits must be non-negative and --max-val must be positive.") | |
| config_path = args.config.resolve() | |
| config = json.loads(config_path.read_text(encoding="utf-8")) | |
| config_base = config_path.parent.parent if config_path.parent.name == "configs" else config_path.parent | |
| output_root = args.output_dir.resolve() if args.output_dir else None | |
| changed = False | |
| for task_index, task in enumerate(config.get("tasks", [])): | |
| task_name = str(task.get("task_name", "task")) | |
| splits = task.get("splits", {}) | |
| if not isinstance(splits, dict) or not splits.get("train"): | |
| print(f"SKIP {task_name}: no train split") | |
| continue | |
| train_path = _resolve_manifest_path(str(splits["train"]), config_base) | |
| records, shape = read_records(train_path) | |
| train_records, val_records = split_records( | |
| records, | |
| val_ratio=args.val_ratio, | |
| min_val=args.min_val, | |
| max_val=args.max_val, | |
| seed=args.seed + task_index, | |
| ) | |
| target_dir = output_root / task_name if output_root else train_path.parent | |
| suffix = ".jsonl" if shape == "jsonl" else ".json" | |
| split_train_path = target_dir / f"train_split{suffix}" | |
| split_val_path = target_dir / f"validation_split{suffix}" | |
| print( | |
| f"{task_name}: {len(records)} -> train={len(train_records)} " | |
| f"validation={len(val_records)} ({split_train_path}, {split_val_path})" | |
| ) | |
| if args.dry_run: | |
| continue | |
| if not args.overwrite and (split_train_path.exists() or split_val_path.exists()): | |
| raise FileExistsError( | |
| f"Split file already exists for {task_name}; use --overwrite to replace." | |
| ) | |
| write_records(split_train_path, train_records, shape) | |
| write_records(split_val_path, val_records, shape) | |
| splits["train"] = _portable_path(split_train_path, config_base) | |
| splits["validation"] = _portable_path(split_val_path, config_base) | |
| splits["val"] = splits["validation"] | |
| splits.setdefault("anchor", splits["validation"]) | |
| changed = True | |
| if args.output_config and changed and not args.dry_run: | |
| args.output_config.parent.mkdir(parents=True, exist_ok=True) | |
| args.output_config.write_text( | |
| json.dumps(config, ensure_ascii=False, indent=2) + "\n", encoding="utf-8" | |
| ) | |
| elif changed and not args.dry_run: | |
| print("Split files written. Pass --output-config to write an updated task-stream config.") | |
| return 0 | |
| def _resolve_manifest_path(raw_path: str, config_base: Path) -> Path: | |
| path = Path(raw_path).expanduser() | |
| return path.resolve() if path.is_absolute() else (config_base / path).resolve() | |
| def _portable_path(path: Path, config_base: Path) -> str: | |
| try: | |
| return str(path.resolve().relative_to(config_base.resolve())) | |
| except ValueError: | |
| return str(path.resolve()) | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) | |