trl-active-groups / scripts /create_demo.py
lewtun's picture
lewtun HF Staff
Refresh group fraction demo and publish the reproducible aggregate inputs
0e317b2 verified
Raw History Blame Contribute Delete
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()