Self-Forcing-part-2 / scripts /merge_predictor_v4_manifests.py
Cccccz's picture
Add files using upload-large-folder tool
2847d0b verified
Raw
History Blame Contribute Delete
5.81 kB
#!/usr/bin/env python3
"""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()