Skip to content

Commit 2476276

Browse files
authored
Merge pull request #45 from interscript/feat/publish-variant-key
publish: --key for precision-variant index entries
2 parents dcb708d + e60a3f4 commit 2476276

1 file changed

Lines changed: 18 additions & 13 deletions

File tree

‎scripts/publish_model.py‎

Lines changed: 18 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -143,6 +143,10 @@ def main() -> None:
143143
parser.add_argument("model_id")
144144
parser.add_argument("--zip", type=Path, required=True)
145145
parser.add_argument("--repo", default=DEFAULT_REPO)
146+
parser.add_argument(
147+
"--key",
148+
help="models.yaml key when publishing a precision variant of an "
149+
"existing base id, e.g. --key <base>-int8 (default: model_id)")
146150
args = parser.parse_args()
147151

148152
# git and gh must operate on this repo regardless of the caller's cwd
@@ -159,6 +163,7 @@ def main() -> None:
159163
meta = load_metadata(args.zip)
160164
if meta["id"] != args.model_id:
161165
raise SystemExit(f"metadata id {meta['id']!r} != requested {args.model_id!r}")
166+
key = args.key or args.model_id
162167

163168
whole_sha = sha256_file(args.zip)
164169
size = args.zip.stat().st_size
@@ -171,7 +176,7 @@ def main() -> None:
171176

172177
tag = args.model_id
173178
with tempfile.NamedTemporaryFile("w", suffix=".md", delete=False) as fh:
174-
fh.write(release_notes(args.model_id, meta, args.zip.name, size, whole_sha, assets))
179+
fh.write(release_notes(key, meta, args.zip.name, size, whole_sha, assets))
175180
notes_path = fh.name
176181

177182
existing = subprocess.run(
@@ -206,8 +211,8 @@ def main() -> None:
206211
# Branch work happens in a dedicated worktree so the caller's tree is
207212
# never checked out (dirty files must not block publication, and
208213
# publication must not clobber in-progress edits).
209-
branch = f"release/{args.model_id}"
210-
worktree = REPO_ROOT.parent / f".wt-publish-{args.model_id}"
214+
branch = f"release/{key}"
215+
worktree = REPO_ROOT.parent / f".wt-publish-{key}"
211216
branches = run(["git", "branch", "--list", branch],
212217
capture_output=True, text=True).stdout
213218
base = branch if branches else "origin/main"
@@ -217,22 +222,22 @@ def main() -> None:
217222
wt_models = worktree / "models.yaml"
218223
wt_repo = str(worktree)
219224
# models/<family>/<id>.metadata.yaml, e.g. models/heb-diac/heb-diac-1.0...
220-
family = args.model_id.rsplit("-", 1)[0]
221-
upsert_models_yaml(wt_models, args.model_id,
222-
entry_block(args.model_id, meta, args.zip.name, whole_sha,
225+
family = key.rsplit("-", 1)[0]
226+
upsert_models_yaml(wt_models, key,
227+
entry_block(key, meta, args.zip.name, whole_sha,
223228
size, assets, args.repo, tag))
224229
model_dir = worktree / "models" / family
225230
model_dir.mkdir(parents=True, exist_ok=True)
226-
(model_dir / f"{args.model_id}.metadata.yaml").write_text(
231+
(model_dir / f"{key}.metadata.yaml").write_text(
227232
yaml.safe_dump(meta, sort_keys=False, allow_unicode=True), encoding="utf-8")
228233

229234
def git_wt(*cmd: str):
230235
return run(["git", "-C", wt_repo, *cmd], capture_output=True, text=True)
231236

232-
git_wt("add", "models.yaml", f"models/{family}/{args.model_id}.metadata.yaml")
237+
git_wt("add", "models.yaml", f"models/{family}/{key}.metadata.yaml")
233238
staged = git_wt("diff", "--cached", "--name-only").stdout.split()
234239
if staged:
235-
git_wt("commit", "-m", f"release: {args.model_id} "
240+
git_wt("commit", "-m", f"release: {key} "
236241
f"({meta['precision']}, parity cer_delta {meta['parity']['cer_delta']}pp "
237242
f"on {meta['parity']['samples']} samples)")
238243
git_wt("push", "-u", "origin", branch)
@@ -242,13 +247,13 @@ def git_wt(*cmd: str):
242247
capture_output=True, text=True).stdout
243248
if not yaml.safe_load(prs):
244249
with tempfile.NamedTemporaryFile("w", suffix=".md", delete=False) as fh:
245-
fh.write(f"## Summary\n- publish {args.model_id} ({meta['precision']}): "
250+
fh.write(f"## Summary\n- publish {key} ({meta['precision']}): "
246251
f"GH Release `{tag}` + models.yaml entry"
247252
f"{' (split parts, GitHub 2GiB cap)' if len(assets) > 1 else ''}\n"
248253
f"- parity cer_delta {meta['parity']['cer_delta']}pp on "
249254
f"{meta['parity']['samples']} samples; strict validator gate passed\n\n"
250255
f"## Test plan\n- [ ] CI green\n- [ ] runtime fetch "
251-
f"`Model.load(\"{args.model_id}\")` resolves and verifies\n")
256+
f"`Model.load(\"{key}\")` resolves and verifies\n")
252257
body_path = fh.name
253258
# the branch push propagates asynchronously; an immediate
254259
# pr create reliably fails with "head branch not found"
@@ -257,7 +262,7 @@ def git_wt(*cmd: str):
257262
for attempt in range(3):
258263
proc = subprocess.run(
259264
["gh", "pr", "create", "-R", args.repo,
260-
"--head", branch, "--title", f"release: {args.model_id}",
265+
"--head", branch, "--title", f"release: {key}",
261266
"--body-file", body_path],
262267
capture_output=True, text=True,
263268
)
@@ -269,7 +274,7 @@ def git_wt(*cmd: str):
269274
raise SystemExit("gh pr create failed after retries")
270275
finally:
271276
run(["git", "worktree", "remove", "--force", str(worktree)])
272-
print(f"published {args.model_id}: release {tag}, branch {branch}")
277+
print(f"published {key}: release {tag}, branch {branch}")
273278

274279

275280
if __name__ == "__main__":

0 commit comments

Comments
 (0)