train: hubckpt -- hubckpt.ensure_repo writes LFS attributes before the first checkpoint push (found by P1)
Browse files- train/hubckpt.py +33 -0
train/hubckpt.py
CHANGED
|
@@ -117,6 +117,39 @@ def verify_checkpoint(repo, dirpath, path_in_repo, token=None, repo_type="datase
|
|
| 117 |
"mismatch": mismatch[:10], "extra": extra[:10], "readback": rb}
|
| 118 |
|
| 119 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 120 |
def push_and_prune(repo, dirpath, path_in_repo, api, repo_type="dataset", token=None,
|
| 121 |
prune=True, min_free_gb=6.0):
|
| 122 |
"""One full §3.13 cycle. Raises if verification fails and pruning was requested -- the local copy is
|
|
|
|
| 117 |
"mismatch": mismatch[:10], "extra": extra[:10], "readback": rb}
|
| 118 |
|
| 119 |
|
| 120 |
+
def ensure_repo(repo, api, token=None, repo_type="dataset"):
|
| 121 |
+
"""Create the checkpoint repo and give it LFS patterns, in that order, before anything is pushed.
|
| 122 |
+
|
| 123 |
+
A 1.6 GB checkpoint is far over the 10 MB non-LFS limit, and a fresh repo has no .gitattributes, so a
|
| 124 |
+
first push fails with "should be tracked by LFS" -- ten minutes into the first checkpoint of a GPU
|
| 125 |
+
session that is being billed while it waits. Found by preflight P1 on CPU, which is the whole point
|
| 126 |
+
of running the cycle with mock bytes before the real run.
|
| 127 |
+
"""
|
| 128 |
+
created = False
|
| 129 |
+
try:
|
| 130 |
+
api.create_repo(repo_id=repo, repo_type=repo_type, exist_ok=True, token=token)
|
| 131 |
+
created = True
|
| 132 |
+
except Exception as e:
|
| 133 |
+
# exist_ok=True should absorb the common case; anything else is worth seeing, not swallowing.
|
| 134 |
+
print(f"ensure_repo create: {type(e).__name__}: {str(e)[:160]}", flush=True)
|
| 135 |
+
have = set()
|
| 136 |
+
try:
|
| 137 |
+
have = set(api.list_repo_files(repo_id=repo, repo_type=repo_type))
|
| 138 |
+
except Exception as e:
|
| 139 |
+
print(f"ensure_repo listing failed ({type(e).__name__}); writing attributes anyway", flush=True)
|
| 140 |
+
if ".gitattributes" not in have:
|
| 141 |
+
body = ("*.safetensors filter=lfs diff=lfs merge=lfs -text\n"
|
| 142 |
+
"*.bin filter=lfs diff=lfs merge=lfs -text\n"
|
| 143 |
+
"*.pt filter=lfs diff=lfs merge=lfs -text\n"
|
| 144 |
+
"*.npy filter=lfs diff=lfs merge=lfs -text\n"
|
| 145 |
+
"*.ckpt filter=lfs diff=lfs merge=lfs -text\n")
|
| 146 |
+
api.upload_file(path_or_fileobj=body.encode(), path_in_repo=".gitattributes", repo_id=repo,
|
| 147 |
+
repo_type=repo_type, commit_message="track checkpoint tensors with LFS",
|
| 148 |
+
token=token)
|
| 149 |
+
return {"repo": repo, "created": created, "had_gitattributes": ".gitattributes" in have,
|
| 150 |
+
"files_before": len(have)}
|
| 151 |
+
|
| 152 |
+
|
| 153 |
def push_and_prune(repo, dirpath, path_in_repo, api, repo_type="dataset", token=None,
|
| 154 |
prune=True, min_free_gb=6.0):
|
| 155 |
"""One full §3.13 cycle. Raises if verification fails and pruning was requested -- the local copy is
|