| |
| """Merge worker shards and create the chunk-1..6 Predictor training manifest.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| import os |
| from pathlib import Path |
| from typing import Any |
|
|
|
|
| ROOT = Path(__file__).resolve().parents[1] |
| DEFAULT_ROOT = Path( |
| "/mnt/local_nvme/zoubin/cz/self_forcing_predictor_v4_1000_seed0" |
| ) |
| NUM_CHUNKS = 7 |
| CHUNK_FRAMES = 3 |
|
|
|
|
| def read_jsonl(path: Path) -> list[dict[str, Any]]: |
| result = [] |
| with path.open("r", encoding="utf-8") as handle: |
| for line_number, line in enumerate(handle, start=1): |
| if not line.strip(): |
| continue |
| try: |
| result.append(json.loads(line)) |
| except json.JSONDecodeError as exc: |
| raise ValueError(f"invalid JSON at {path}:{line_number}") from exc |
| return result |
|
|
|
|
| def atomic_write(path: Path, text: str) -> None: |
| path.parent.mkdir(parents=True, exist_ok=True) |
| temporary = path.with_name(f".{path.name}.tmp.{os.getpid()}") |
| with temporary.open("w", encoding="utf-8") as handle: |
| handle.write(text) |
| handle.flush() |
| os.fsync(handle.fileno()) |
| os.replace(temporary, path) |
|
|
|
|
| def jsonl_text(items: list[dict[str, Any]]) -> str: |
| return "".join( |
| json.dumps(item, ensure_ascii=False, sort_keys=True) + "\n" |
| for item in items |
| ) |
|
|
|
|
| def main() -> None: |
| parser = argparse.ArgumentParser(description=__doc__) |
| parser.add_argument("--dataset_root", type=Path, default=DEFAULT_ROOT) |
| parser.add_argument("--num_workers", type=int, default=8) |
| parser.add_argument( |
| "--allow_incomplete", |
| action="store_true", |
| help="merge complete case prefixes for inspection; formal training must not use this", |
| ) |
| args = parser.parse_args() |
| root = args.dataset_root.resolve() |
| cases_path = root / "cases.jsonl" |
| if not cases_path.is_file(): |
| raise FileNotFoundError(cases_path) |
| cases = read_jsonl(cases_path) |
| expected_case_ids = {int(item["case_id"]) for item in cases} |
| records: dict[tuple[int, int], dict[str, Any]] = {} |
| for worker_id in range(args.num_workers): |
| path = root / "manifests" / f"worker_{worker_id:02d}.jsonl" |
| if not path.is_file(): |
| if args.allow_incomplete: |
| continue |
| raise FileNotFoundError(path) |
| for item in read_jsonl(path): |
| case_id = int(item["case_id"]) |
| chunk_id = int(item["chunk_id"]) |
| if case_id % args.num_workers != worker_id: |
| raise ValueError( |
| f"case {case_id} is in worker {worker_id}, expected " |
| f"worker {case_id % args.num_workers}" |
| ) |
| if case_id not in expected_case_ids or not 0 <= chunk_id < NUM_CHUNKS: |
| raise ValueError(f"unexpected record case={case_id}, chunk={chunk_id}") |
| key = (case_id, chunk_id) |
| if key in records and records[key] != item: |
| raise ValueError(f"conflicting duplicate record {key}") |
| records[key] = item |
|
|
| complete_case_ids: list[int] = [] |
| incomplete: dict[int, list[int]] = {} |
| for case_id in sorted(expected_case_ids): |
| missing = [ |
| chunk_id |
| for chunk_id in range(NUM_CHUNKS) |
| if (case_id, chunk_id) not in records |
| ] |
| if missing: |
| incomplete[case_id] = missing |
| else: |
| complete_case_ids.append(case_id) |
| if incomplete and not args.allow_incomplete: |
| preview = list(incomplete.items())[:10] |
| raise RuntimeError( |
| f"{len(incomplete)} cases are incomplete; first missing chunks: {preview}" |
| ) |
|
|
| merged: list[dict[str, Any]] = [] |
| train: list[dict[str, Any]] = [] |
| for case_id in complete_case_ids: |
| history: dict[str, list[str]] = { |
| str(block_id): [] |
| for block_id in records[(case_id, 0)]["clean_prefeature_files"] |
| } |
| for chunk_id in range(NUM_CHUNKS): |
| source = dict(records[(case_id, chunk_id)]) |
| source["previous_step_tensor_file"] = ( |
| None |
| if chunk_id == 0 |
| else records[(case_id, chunk_id - 1)]["step_tensor_file"] |
| ) |
| source["history_clean_prefeature_files"] = { |
| block_id: list(paths) |
| for block_id, paths in history.items() |
| } |
| source["context_frames"] = chunk_id * CHUNK_FRAMES |
| merged.append(source) |
| if chunk_id > 0: |
| train.append(source) |
| for block_id, path in source["clean_prefeature_files"].items(): |
| history.setdefault(str(block_id), []).append(str(path)) |
|
|
| atomic_write(root / "manifest.jsonl", jsonl_text(merged)) |
| atomic_write(root / "train_manifest.jsonl", jsonl_text(train)) |
| summary = { |
| "allow_incomplete": bool(args.allow_incomplete), |
| "expected_cases": len(expected_case_ids), |
| "complete_cases": len(complete_case_ids), |
| "incomplete_cases": len(incomplete), |
| "all_chunk_records": len(merged), |
| "training_chunk_records": len(train), |
| "training_pairs_per_chunk": 3, |
| "training_pairs": len(train) * 3, |
| "chunk0_excluded_from_training": True, |
| "history_storage": ( |
| "incremental clean files; history_clean_prefeature_files lists chunks [0, j)" |
| ), |
| } |
| atomic_write( |
| root / "manifest_summary.json", |
| json.dumps(summary, indent=2, sort_keys=True) + "\n", |
| ) |
| print( |
| f"Merged {len(merged)} chunk records from {len(complete_case_ids)} cases; " |
| f"train_manifest has {len(train)} chunks / {len(train) * 3} adjacent-step pairs." |
| ) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|