Ariadne-Laya-TD / make_training_config.py
GoatHerder's picture
Upload 49 files
5f2a11f verified
Raw History Blame
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")