Download scripts/build_textual_stream.py from RL-MIND/Harness-Continual-Learning: direct link, hf CLI and curl.
- Browser
- Download file 9.85 kB
-
https://huggingface.co/RL-MIND/Harness-Continual-Learning/resolve/main/scripts/build_textual_stream.py
- Command line
-
hf download hf://RL-MIND/Harness-Continual-Learning/scripts/build_textual_stream.py
-
curl -L -o build_textual_stream.py https://huggingface.co/RL-MIND/Harness-Continual-Learning/resolve/main/scripts/build_textual_stream.py
9.85 kB
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import random | |
| from collections import defaultdict | |
| from pathlib import Path | |
| from typing import Any, Callable, Hashable, Iterable | |
| PROJECT_ROOT = Path(__file__).resolve().parents[1] | |
| SEED = 20260718 | |
| def build_parser() -> argparse.ArgumentParser: | |
| parser = argparse.ArgumentParser( | |
| description="Build the four-task textual-reasoning HCL data splits." | |
| ) | |
| parser.add_argument( | |
| "--musique-root", | |
| type=Path, | |
| default=PROJECT_ROOT / "data" / "musique_data" / "data", | |
| help="Directory containing musique_ans_v1.0_{train,dev}.jsonl.", | |
| ) | |
| parser.add_argument( | |
| "--source-root", | |
| type=Path, | |
| default=PROJECT_ROOT / "data" / "source_4task", | |
| help="Directory containing existing gsm8k/proofwriter/hotpotqa split folders.", | |
| ) | |
| parser.add_argument( | |
| "--output-root", | |
| type=Path, | |
| default=PROJECT_ROOT / "data" / "datasets", | |
| help="Destination used by configs/taskstream_textual_main_250_50_500.json.", | |
| ) | |
| parser.add_argument("--train-size", type=int, default=500) | |
| parser.add_argument("--validation-size", type=int, default=100) | |
| parser.add_argument("--anchor-size", type=int, default=100) | |
| parser.add_argument("--test-size", type=int, default=500) | |
| parser.add_argument("--seed", type=int, default=SEED) | |
| return parser | |
| def main() -> int: | |
| args = build_parser().parse_args() | |
| sizes = { | |
| "train": args.train_size, | |
| "validation": args.validation_size, | |
| "anchor": args.anchor_size, | |
| "test": args.test_size, | |
| } | |
| if any(size < 0 for size in sizes.values()): | |
| raise ValueError("Split sizes must be non-negative.") | |
| for task_name in ("gsm8k", "proofwriter", "hotpotqa"): | |
| build_existing_task(args.source_root, args.output_root, task_name, sizes=sizes) | |
| build_musique( | |
| args.musique_root, | |
| args.output_root / "musique", | |
| sizes=sizes, | |
| seed=args.seed, | |
| ) | |
| print(f"Prepared four-task data at {args.output_root.resolve()}") | |
| for task_name in ("gsm8k", "proofwriter", "hotpotqa", "musique"): | |
| counts = { | |
| path.stem: count_jsonl(path) | |
| for path in sorted((args.output_root / task_name).glob("*.jsonl")) | |
| } | |
| print(f"{task_name}: {counts}") | |
| return 0 | |
| def build_existing_task( | |
| source_root: Path, | |
| output_root: Path, | |
| task_name: str, | |
| *, | |
| sizes: dict[str, int], | |
| ) -> None: | |
| source_dir = source_root / task_name | |
| splits = { | |
| split: read_jsonl(source_dir / f"{split}.jsonl")[: sizes[split]] | |
| for split in ("train", "validation", "test") | |
| } | |
| for split, rows in splits.items(): | |
| require_count(f"{task_name} {split}", rows, sizes[split]) | |
| # This matches the recovered experiment protocol for the three existing tasks. | |
| # A distinct MuSiQue anchor sample is constructed below. | |
| require_count( | |
| f"{task_name} validation-as-anchor", | |
| splits["validation"], | |
| sizes["anchor"], | |
| ) | |
| splits["anchor"] = [dict(row) for row in splits["validation"][: sizes["anchor"]]] | |
| write_splits(output_root / task_name, splits) | |
| def build_musique( | |
| source_root: Path, | |
| output_dir: Path, | |
| *, | |
| sizes: dict[str, int], | |
| seed: int, | |
| ) -> None: | |
| train_source = read_jsonl(source_root / "musique_ans_v1.0_train.jsonl") | |
| dev_source = read_jsonl(source_root / "musique_ans_v1.0_dev.jsonl") | |
| train_pool = stratified_shuffle(train_source, key=musique_hop_count, seed=seed + 300) | |
| required = sizes["train"] + sizes["validation"] + sizes["anchor"] | |
| require_count("musique train/validation/anchor", train_pool, required) | |
| raw_splits = { | |
| "train": train_pool[: sizes["train"]], | |
| "validation": train_pool[ | |
| sizes["train"] : sizes["train"] + sizes["validation"] | |
| ], | |
| "anchor": train_pool[ | |
| sizes["train"] + sizes["validation"] : required | |
| ], | |
| "test": stratified_sample( | |
| dev_source, | |
| size=sizes["test"], | |
| key=musique_hop_count, | |
| seed=seed + 301, | |
| ), | |
| } | |
| splits = { | |
| split: [convert_musique(row, split, index) for index, row in enumerate(rows)] | |
| for split, rows in raw_splits.items() | |
| } | |
| assert_disjoint_source_ids(splits) | |
| write_splits(output_dir, splits) | |
| def convert_musique(row: dict[str, Any], split: str, index: int) -> dict[str, Any]: | |
| answer = str(row.get("answer", "")).strip() | |
| if not answer: | |
| raise ValueError(f"MuSiQue row has no answer: {row.get('id')}") | |
| context = [ | |
| { | |
| "idx": paragraph.get("idx"), | |
| "title": str(paragraph.get("title", "")), | |
| "paragraph_text": str(paragraph.get("paragraph_text", "")), | |
| } | |
| for paragraph in row.get("paragraphs", []) | |
| ] | |
| if not context: | |
| raise ValueError(f"MuSiQue row has no paragraphs: {row.get('id')}") | |
| return { | |
| "id": f"musique_{split}_{index:05d}", | |
| "question": ( | |
| f"{str(row.get('question', '')).strip()}\n\n" | |
| "Answer using the provided context. Return only the shortest final answer; do not explain." | |
| ), | |
| "visible_context": context, | |
| "answer": answer, | |
| "answers": unique_nonempty([answer, *row.get("answer_aliases", [])]), | |
| "task_name": "musique", | |
| "task_type": "compositional_multi_hop_qa", | |
| "metadata": { | |
| "source": "musique", | |
| "original_id": row.get("id"), | |
| "hop_count": musique_hop_count(row), | |
| "paragraph_count": len(context), | |
| }, | |
| } | |
| def musique_hop_count(row: dict[str, Any]) -> int: | |
| identifier = str(row.get("id", "")) | |
| try: | |
| return int(identifier.split("hop", 1)[0]) | |
| except ValueError: | |
| decomposition = row.get("question_decomposition", []) | |
| return len(decomposition) if isinstance(decomposition, list) else 0 | |
| def stratified_sample( | |
| rows: list[dict[str, Any]], | |
| *, | |
| size: int, | |
| key: Callable[[dict[str, Any]], Hashable], | |
| seed: int, | |
| ) -> list[dict[str, Any]]: | |
| if size > len(rows): | |
| raise ValueError(f"Cannot sample {size} rows from a population of {len(rows)}") | |
| grouped = _shuffled_groups(rows, key=key, seed=seed) | |
| exact = {name: size * len(group) / len(rows) for name, group in grouped.items()} | |
| allocation = {name: int(value) for name, value in exact.items()} | |
| remaining = size - sum(allocation.values()) | |
| remainder_order = sorted( | |
| grouped, | |
| key=lambda name: (exact[name] - allocation[name], str(name)), | |
| reverse=True, | |
| ) | |
| for name in remainder_order[:remaining]: | |
| allocation[name] += 1 | |
| selected = [ | |
| row | |
| for name in sorted(grouped, key=str) | |
| for row in grouped[name][: allocation[name]] | |
| ] | |
| random.Random(seed).shuffle(selected) | |
| return selected | |
| def stratified_shuffle( | |
| rows: list[dict[str, Any]], | |
| *, | |
| key: Callable[[dict[str, Any]], Hashable], | |
| seed: int, | |
| ) -> list[dict[str, Any]]: | |
| """Interleave shuffled strata so contiguous slices retain every stratum.""" | |
| grouped = _shuffled_groups(rows, key=key, seed=seed) | |
| consumed = {name: 0 for name in grouped} | |
| output: list[dict[str, Any]] = [] | |
| while len(output) < len(rows): | |
| candidates = [name for name in grouped if consumed[name] < len(grouped[name])] | |
| name = max( | |
| candidates, | |
| key=lambda candidate: ( | |
| len(grouped[candidate]) / len(rows) | |
| - consumed[candidate] / max(len(output), 1), | |
| -consumed[candidate], | |
| str(candidate), | |
| ), | |
| ) | |
| output.append(grouped[name][consumed[name]]) | |
| consumed[name] += 1 | |
| return output | |
| def _shuffled_groups( | |
| rows: list[dict[str, Any]], | |
| *, | |
| key: Callable[[dict[str, Any]], Hashable], | |
| seed: int, | |
| ) -> dict[Hashable, list[dict[str, Any]]]: | |
| grouped: dict[Hashable, list[dict[str, Any]]] = defaultdict(list) | |
| for row in rows: | |
| grouped[key(row)].append(row) | |
| rng = random.Random(seed) | |
| for group in grouped.values(): | |
| rng.shuffle(group) | |
| return grouped | |
| def assert_disjoint_source_ids(splits: dict[str, list[dict[str, Any]]]) -> None: | |
| seen: dict[Any, str] = {} | |
| for split, rows in splits.items(): | |
| for row in rows: | |
| identifier = row["metadata"]["original_id"] | |
| if identifier in seen: | |
| raise ValueError( | |
| f"MuSiQue source id {identifier!r} appears in both {seen[identifier]} and {split}" | |
| ) | |
| seen[identifier] = split | |
| def unique_nonempty(values: Iterable[Any]) -> list[str]: | |
| output: list[str] = [] | |
| for value in values: | |
| text = str(value).strip() | |
| if text and text not in output: | |
| output.append(text) | |
| return output | |
| def read_jsonl(path: Path) -> list[dict[str, Any]]: | |
| with path.open("r", encoding="utf-8") as handle: | |
| return [json.loads(line) for line in handle if line.strip()] | |
| def write_splits(output_dir: Path, splits: dict[str, list[dict[str, Any]]]) -> None: | |
| output_dir.mkdir(parents=True, exist_ok=True) | |
| for split, rows in splits.items(): | |
| with (output_dir / f"{split}.jsonl").open("w", encoding="utf-8") as handle: | |
| for row in rows: | |
| handle.write(json.dumps(row, ensure_ascii=False) + "\n") | |
| def require_count(label: str, rows: list[Any], expected: int) -> None: | |
| if len(rows) < expected: | |
| raise ValueError(f"{label} has {len(rows)} rows, but {expected} are required") | |
| def count_jsonl(path: Path) -> int: | |
| with path.open("r", encoding="utf-8") as handle: | |
| return sum(1 for line in handle if line.strip()) | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) | |