File size: 5,660 Bytes
ffdcfe7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
"""Stage and upload the campaign, then verify it reads back before anything is destroyed.

/workspace is not a volume on this instance, so the container filesystem is the
only copy until this runs. This project has already lost two batches exactly that
way -- the olmix-policy/ft-as20k runs behind policy-mixture.md, and 256 raw
validation dumps from ranking-transfer.md -- so the verify step is not optional
and the instance is not destroyed on an unverified upload.

Uploads every measurement and provenance record plus the final checkpoint of each
run. Intermediate depth checkpoints (14 per run, ~100 GB) are deliberately left
behind: their probe scores are archived, and they are reproducible from the
recipe, the pinned mixture vector and the seed.
"""
import argparse
import hashlib
import json
import shutil
from pathlib import Path

from huggingface_hub import HfApi

STAGE = Path("/workspace/upload")
RUNS = Path("/workspace/runs")
KEEP = ("run.json", "config.yaml", "status.json", "events.jsonl")


def sha(p: Path, n=1 << 20) -> str:
    h = hashlib.sha256()
    with open(p, "rb") as f:
        while chunk := f.read(n):
            h.update(chunk)
    return h.hexdigest()


def stage() -> dict:
    if STAGE.exists():
        shutil.rmtree(STAGE)
    manifest = {}
    for campaign in ("mae-64", "mae-seed", "mae-lr", "mae-validate"):
        root = RUNS / campaign
        if not root.exists():
            continue
        for run in sorted(root.iterdir()):
            if not run.is_dir():
                continue
            dest = STAGE / campaign / run.name
            dest.mkdir(parents=True, exist_ok=True)
            for name in KEEP:
                if (run / name).exists():
                    shutil.copy2(run / name, dest / name)
            exports = sorted((run / "exports").glob("step_*"))
            for exp in exports:
                # every readout at every depth, but weights only at the last
                for j in list(exp.glob("probe.json")) + list(exp.glob("lora*.json")):
                    d = dest / "exports" / exp.name
                    d.mkdir(parents=True, exist_ok=True)
                    shutil.copy2(j, d / j.name)
            if exports:
                d = dest / "exports" / exports[-1].name
                d.mkdir(parents=True, exist_ok=True)
                for name in ("model.safetensors", "config.yaml"):
                    if (exports[-1] / name).exists():
                        shutil.copy2(exports[-1] / name, d / name)
    # analysis outputs, campaign specs, and the code that produced them
    for src, sub in ((Path("/workspace/analysis"), "analysis"),
                     (Path("/workspace/scripts"), "scripts"),
                     (Path("/workspace/configs"), "configs")):
        for f in sorted(src.rglob("*")):
            if f.is_file() and f.suffix in (".json", ".py", ".yaml", ".csv"):
                d = STAGE / sub / f.relative_to(src)
                d.parent.mkdir(parents=True, exist_ok=True)
                shutil.copy2(f, d)
    shutil.copy2("/workspace/code/eat-map-regmix/eatmap/mae.py", STAGE / "scripts" / "mae.py")
    for f in sorted(STAGE.rglob("*")):
        if f.is_file():
            manifest[str(f.relative_to(STAGE))] = {"size": f.stat().st_size, "sha256": sha(f)}
    (STAGE / "MANIFEST.json").write_text(json.dumps(manifest, indent=1, sort_keys=True))
    return manifest


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--repo", default="quinnlue/mae-cross-objective")
    ap.add_argument("--private", action="store_true")
    ap.add_argument("--verify-only", action="store_true")
    args = ap.parse_args()

    api = HfApi()
    if not args.verify_only:
        m = stage()
        total = sum(v["size"] for v in m.values())
        print(f"staged {len(m)} files, {total/2**30:.2f} GiB -> {args.repo}")
        api.create_repo(args.repo, repo_type="model", private=args.private, exist_ok=True)
        api.upload_folder(folder_path=str(STAGE), repo_id=args.repo, repo_type="model",
                          commit_message="MAE cross-objective mixture-ranking campaign")
        print("upload complete")

    # ---- verify: every staged file present on the Hub at the right size, and a
    # random sample byte-identical after a round trip.
    manifest = json.loads((STAGE / "MANIFEST.json").read_text())
    remote = {f.rfilename: f for f in api.repo_info(
        args.repo, repo_type="model", files_metadata=True).siblings}
    missing = [k for k in manifest if k not in remote]
    wrong = [k for k, v in manifest.items()
             if k in remote and remote[k].size not in (None, v["size"])]
    print(f"\nverify: {len(manifest)} staged | {len(remote)} on hub | "
          f"missing {len(missing)} | size-mismatch {len(wrong)}")
    if missing:
        print("  MISSING:", missing[:10])
    if wrong:
        print("  MISMATCH:", wrong[:10])

    import random
    from huggingface_hub import hf_hub_download
    random.seed(0)
    sample = random.sample([k for k in manifest if k in remote],
                           min(8, len(manifest)))
    bad = []
    for k in sample:
        p = hf_hub_download(args.repo, k, repo_type="model",
                            cache_dir="/workspace/.verify")
        if sha(Path(p)) != manifest[k]["sha256"]:
            bad.append(k)
    print(f"round-trip sha256 on {len(sample)} sampled files: "
          f"{'ALL MATCH' if not bad else f'MISMATCH {bad}'}")
    ok = not missing and not wrong and not bad
    print(f"\nARCHIVE {'VERIFIED' if ok else 'NOT VERIFIED -- DO NOT DESTROY'}")
    return 0 if ok else 1


if __name__ == "__main__":
    raise SystemExit(main())