Spaces:
Running
Running
Download scripts/create_demo.py from lewtun/trl-active-groups: direct link, hf CLI and curl.
- Browser
- Download file 4.88 kB
-
https://huggingface.co/spaces/lewtun/trl-active-groups/resolve/main/scripts/create_demo.py
- Command line
-
hf download hf://spaces/lewtun/trl-active-groups/scripts/create_demo.py
-
curl -L -o create_demo.py https://huggingface.co/spaces/lewtun/trl-active-groups/resolve/main/scripts/create_demo.py
4.88 kB
| """Publish the aggregate v17.01 replay as a public Trackio snapshot.""" | |
| import argparse | |
| import csv | |
| import json | |
| import os | |
| import shutil | |
| import tempfile | |
| from pathlib import Path | |
| from huggingface_hub import CommitOperationAdd, HfApi | |
| PROJECT = "trl-active-groups" | |
| RUN = "v17.01-replay" | |
| METRICS = [ | |
| "train/frac_active_groups", | |
| "train/frac_reward_zero_std", | |
| "train/frac_reward_all_zero", | |
| "train/frac_reward_all_one", | |
| "train/reward", | |
| "train/reward_std", | |
| ] | |
| COUNT_KEYS = [ | |
| "scorable_groups", | |
| "active_groups", | |
| "zero_std_groups", | |
| "all_zero_groups", | |
| "all_one_groups", | |
| "unscorable_groups", | |
| ] | |
| def main(): | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--space-id", required=True, help="Your public Hugging Face Space ID.") | |
| args = parser.parse_args() | |
| root = Path(__file__).resolve().parents[1] | |
| provenance = json.loads((root / "data/provenance.json").read_text()) | |
| with (root / "data/v17.01-intervals.csv").open() as handle: | |
| intervals = list(csv.DictReader(handle)) | |
| with tempfile.TemporaryDirectory(prefix="trl-active-groups-") as tmp: | |
| os.environ["TRACKIO_DIR"] = str(Path(tmp) / "trackio") | |
| os.environ["TRACKIO_STORAGE_MODE"] = "sqlite" | |
| os.environ["TRACKIO_PLOT_ORDER"] = ",".join(METRICS) | |
| os.environ.pop("TRACKIO_SPACE_ID", None) | |
| os.environ.pop("TRACKIO_SERVER_URL", None) | |
| import trackio | |
| from trackio.frontend_config import BUNDLED_FRONTEND_DIR | |
| trackio.init( | |
| project=PROJECT, | |
| name=RUN, | |
| group="Real rewards by generation model version", | |
| config=provenance, | |
| auto_log_gpu=False, | |
| auto_log_cpu=False, | |
| embed=False, | |
| ) | |
| for row in intervals: | |
| counts = {key: int(row[key]) for key in COUNT_KEYS} | |
| n = counts["scorable_groups"] | |
| assert n > 0 | |
| assert counts["active_groups"] + counts["zero_std_groups"] == n | |
| assert counts["all_zero_groups"] + counts["all_one_groups"] == counts["zero_std_groups"] | |
| metrics = {f"train/{key}": value for key, value in counts.items()} | |
| metrics.update( | |
| { | |
| "train/frac_active_groups": counts["active_groups"] / n, | |
| "train/frac_reward_zero_std": counts["zero_std_groups"] / n, | |
| "train/frac_reward_all_zero": counts["all_zero_groups"] / n, | |
| "train/frac_reward_all_one": counts["all_one_groups"] / n, | |
| "train/reward": float(row["reward"]), | |
| "train/reward_std": float(row["reward_std"]), | |
| } | |
| ) | |
| trackio.log(metrics, step=int(row["model_version"])) | |
| trackio.finish() | |
| frontend = Path(tmp) / "frontend" | |
| shutil.copytree(BUNDLED_FRONTEND_DIR, frontend) | |
| index = frontend / "index.html" | |
| html = index.read_text().replace("Trackio Dashboard", "TRL active group metrics") | |
| banner = ( | |
| '<div style="padding:12px 20px;background:#e8f5f1;color:#174a40;' | |
| 'font:14px system-ui;border-bottom:1px solid #b5d9cf">' | |
| "<strong>v17.01 reward replay</strong> 路 101,760 real groups 路 " | |
| "grouped by generation model version 路 " | |
| '<a href="https://huggingface.co/spaces/' + args.space_id + '/blob/main/README.md" ' | |
| 'target="_blank" rel="noopener">Methodology</a> 路 ' | |
| '<a href="https://github.com/huggingface/trl-internal/pull/317" ' | |
| 'target="_blank" rel="noopener">PR #317</a></div>' | |
| ) | |
| defaults = ( | |
| "<script>const u=new URL(location.href);" | |
| 'if(!u.searchParams.has("metrics")){' | |
| 'u.searchParams.set("project",' + json.dumps(PROJECT) + ");" | |
| 'u.searchParams.set("metrics",' + json.dumps(",".join(METRICS)) + ");" | |
| 'history.replaceState(null,"",u);}</script>' | |
| ) | |
| index.write_text(html.replace("<body>", "<body>" + banner + defaults)) | |
| trackio.sync( | |
| project=PROJECT, | |
| space_id=args.space_id, | |
| bucket_id=f"{args.space_id}-bucket", | |
| sdk="static", | |
| private=False, | |
| frontend_dir=frontend, | |
| ) | |
| paths = [root / "README.md", Path(__file__), *sorted((root / "data").glob("*"))] | |
| HfApi().create_commit( | |
| repo_id=args.space_id, | |
| repo_type="space", | |
| operations=[ | |
| CommitOperationAdd(path_in_repo=path.relative_to(root).as_posix(), path_or_fileobj=str(path)) | |
| for path in paths | |
| ], | |
| commit_message="Refresh group fraction demo and publish the reproducible aggregate inputs", | |
| ) | |
| print(f"https://huggingface.co/spaces/{args.space_id}") | |
| if __name__ == "__main__": | |
| main() | |