File size: 2,165 Bytes
5ccb4fd | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 | # Copyright (c) 2026 Simulacra Research Inc.
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import argparse
import dataclasses
import json
from pathlib import Path
from .config import load_config
def _write_metric(path: Path, metric) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("a", encoding="utf-8") as stream:
stream.write(
json.dumps(dataclasses.asdict(metric), separators=(",", ":")) + "\n"
)
def _parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(prog="hamiltonzero")
commands = parser.add_subparsers(dest="mode", required=True)
for mode in ("train", "finetune"):
command = commands.add_parser(mode)
command.add_argument("config", type=Path)
command.add_argument("--reuse-mcmc", type=Path)
evaluate = commands.add_parser("eval")
evaluate.add_argument("config", type=Path)
pathway = evaluate.add_mutually_exclusive_group()
pathway.add_argument("--contest", action="store_true")
pathway.add_argument("--large-n", action="store_true")
return parser
def main(argv: list[str] | None = None) -> None:
args = _parser().parse_args(argv)
config = load_config(args.config, args.mode)
if args.mode == "eval":
from .modes.eval import run
if args.contest or args.large_n:
config = dataclasses.replace(
config,
contest=bool(args.contest),
large_n=bool(args.large_n),
)
run(config)
return
if args.reuse_mcmc is not None:
config = dataclasses.replace(
config,
mcmc=dataclasses.replace(config.mcmc, reuse_mcmc=args.reuse_mcmc),
)
metrics_path = config.output.with_name(config.output.name + ".metrics.jsonl")
sink = lambda metric: _write_metric(metrics_path, metric)
if args.mode == "train":
from .modes.train import run_train
run_train(config, metric_sink=sink)
else:
from .modes.finetune import run_finetune
run_finetune(config, metric_sink=sink)
if __name__ == "__main__":
main()
|