File size: 1,955 Bytes
17d5066 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 | #!/usr/bin/env python3
"""Build sealed MNIST mask banks or registered C5 long-path shards."""
from __future__ import annotations
import argparse
import json
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
from repro_control.data import (
build_mnist_mask_bank,
generate_c5_size,
write_mnist_mask_bank,
)
from repro_control.hashing import atomic_write_json
def main() -> int:
parser = argparse.ArgumentParser()
sub = parser.add_subparsers(dest="command", required=True)
masks = sub.add_parser("mnist-masks")
masks.add_argument("--dataset-revision", required=True)
masks.add_argument("--example-ids-json", type=Path, required=True)
masks.add_argument("--output", type=Path, required=True)
c5 = sub.add_parser("c5")
c5.add_argument("--size", type=int, required=True)
c5.add_argument("--training-identities-json", type=Path, required=True)
c5.add_argument("--examples", type=int, default=1000)
c5.add_argument("--output-npz", type=Path, required=True)
c5.add_argument("--output-manifest", type=Path, required=True)
args = parser.parse_args()
if args.command == "mnist-masks":
ids = json.loads(args.example_ids_json.read_text())
bank = build_mnist_mask_bank(args.dataset_revision, ids)
write_mnist_mask_bank(args.output, bank)
print(bank["bank_sha256"])
return 0
import numpy as np
identities = frozenset(json.loads(args.training_identities_json.read_text()))
arrays, manifest = generate_c5_size(
args.size,
training_identities=identities,
examples=args.examples,
)
args.output_npz.parent.mkdir(parents=True, exist_ok=True)
np.savez_compressed(args.output_npz, **arrays)
atomic_write_json(args.output_manifest, manifest)
print(manifest["identity_root_sha256"])
return 0
if __name__ == "__main__":
raise SystemExit(main())
|