Download make_training_config.py from GoatHerder/Ariadne-Laya-TD: direct link, hf CLI and curl.
- Browser
- Download file 1.31 kB
-
https://huggingface.co/GoatHerder/Ariadne-Laya-TD/resolve/5f2a11fa55aa19cfeb212afb945ee83dadd722e9/make_training_config.py
- Command line
-
hf download hf://GoatHerder/Ariadne-Laya-TD@5f2a11fa55aa19cfeb212afb945ee83dadd722e9/make_training_config.py
-
curl -L -o make_training_config.py https://huggingface.co/GoatHerder/Ariadne-Laya-TD/resolve/5f2a11fa55aa19cfeb212afb945ee83dadd722e9/make_training_config.py
1.31 kB
| """Resolve the pinned base checkpoint and create a reproducible training config.""" | |
| import argparse | |
| import json | |
| from pathlib import Path | |
| from huggingface_hub import snapshot_download | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--seed", type=int, choices=range(5), required=True) | |
| parser.add_argument("--device", default="cuda:0") | |
| parser.add_argument("--out", required=True) | |
| args = parser.parse_args() | |
| root = Path(__file__).resolve().parent | |
| metadata = json.loads((root / "adapter_config.json").read_text()) | |
| base = snapshot_download( | |
| metadata["base_model"], | |
| revision=metadata["base_revision"], | |
| allow_patterns=["model.safetensors", "rl_agent_config.json", "encoder/*", "tokenizer/*"], | |
| ) | |
| config = { | |
| "name": f"Deterministic frozen base + linear, seed {args.seed} (UNTRAINED)", | |
| "adapter": "python", | |
| "mode": "reference", | |
| "factory": "ariadne_bench.frozen_input_interface:create", | |
| "options": { | |
| "laya_model": str(Path(base).resolve()), | |
| "device": args.device, | |
| "batch_size": 16, | |
| "max_len": 1024, | |
| "head_max_len": 256, | |
| "seed": args.seed, | |
| "training_scope": "bridge", | |
| "add_interface": True, | |
| "deterministic": True, | |
| }, | |
| } | |
| Path(args.out).write_text(json.dumps(config, indent=2) + "\n") | |