Harness-Continual-Learning / scripts /split_task_stream.py
hhh134's picture
Upload HCL open-source release (part 3)
087e4a0 verified
Raw History Blame Contribute Delete
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())