pactbench / pact /export_foundation_row_keys.py
BBoran's picture
Publish current portable PACTBench release
f1fc3a0 verified
Raw
History Blame Contribute Delete
1.82 kB
#!/usr/bin/env python3
"""Export exact training-row identities from a foundation feature cache."""
from __future__ import annotations
import argparse
import hashlib
import json
import os
import torch
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--feature_cache", required=True)
parser.add_argument("--out", required=True)
args = parser.parse_args()
payload = torch.load(args.feature_cache, map_location="cpu", weights_only=False, mmap=True)
rows = payload.get("rows") or []
keys = [[str(row.get("boss")), int(row.get("fight", -1)), int(row.get("index", -1))] for row in rows]
if not keys:
raise ValueError(f"feature cache contains no rows: {args.feature_cache}")
if len({tuple(key) for key in keys}) != len(keys):
raise ValueError(f"feature cache contains duplicate row keys: {args.feature_cache}")
digest = hashlib.sha256(
json.dumps(keys, ensure_ascii=False, separators=(",", ":")).encode("utf-8")
).hexdigest()
value = {
"schema_revision": "pact-matched-training-rows-v1",
"source_feature_cache": os.path.realpath(args.feature_cache),
"source_request_signature": payload.get("request_signature"),
"n_rows": len(keys),
"row_keys_sha256": digest,
"row_keys": keys,
}
os.makedirs(os.path.dirname(args.out) or ".", exist_ok=True)
tmp = f"{args.out}.tmp.{os.getpid()}"
with open(tmp, "w", encoding="utf-8") as handle:
json.dump(value, handle, ensure_ascii=False, indent=2, allow_nan=False)
handle.flush()
os.fsync(handle.fileno())
os.replace(tmp, args.out)
print(json.dumps({key: value[key] for key in value if key != "row_keys"}, indent=2))
print(f"wrote {args.out}")
if __name__ == "__main__":
main()