"""Command-line interface for generation and all offline checks.""" from __future__ import annotations import argparse import asyncio import json from pathlib import Path from typing import Any from .clients import OpenAIJsonClient from .models import CanonicalItem from .offline import ( align_run_to_release, load_dataset, summarize_run, validate_public_dataset, ) from .pipeline import KernelPipeline, PipelineConfig from .release import export_release from .store import RunStore from .surface import SurfacePipeline def write_or_print(payload: dict[str, Any], output: Path | None) -> None: text = json.dumps(payload, ensure_ascii=False, indent=2, sort_keys=True) + "\n" if output is None: print(text, end="") else: output.parent.mkdir(parents=True, exist_ok=True) output.write_text(text, encoding="utf-8") async def generate_kernel(args: argparse.Namespace) -> None: records = load_dataset(args.dataset) if args.item_id not in records: raise SystemExit(f"item ID {args.item_id!r} not found in {args.dataset}") item = CanonicalItem.from_public_record(records[args.item_id]) config = PipelineConfig( proposer_model=args.proposer_model, judge_model=args.judge_model, ) proposer = OpenAIJsonClient(args.proposer_model) judges = [ OpenAIJsonClient( args.judge_model, seed=None if args.seed is None else args.seed + judge_id, ) for judge_id in range(1, 6) ] pipeline = KernelPipeline( proposer=proposer, judges=judges, store=RunStore(args.run_dir, item.item_id), config=config, ) result = await pipeline.run(item) print(result.model_dump_json(indent=2)) async def generate_surfaces(args: argparse.Namespace) -> None: records = load_dataset(args.dataset) if args.item_id not in records: raise SystemExit(f"item ID {args.item_id!r} not found in {args.dataset}") item = CanonicalItem.from_public_record(records[args.item_id]) store = RunStore(args.run_dir, item.item_id) store.write_input(item) store.write_config( { "protocol_name": "gap-surface-v1", "proposer_model": args.proposer_model, } ) pipeline = SurfacePipeline(OpenAIJsonClient(args.proposer_model), store) variants = await pipeline.run_all(item) print( json.dumps( { family: variant.as_release_payload() for family, variant in variants.items() }, ensure_ascii=False, indent=2, ) ) def build_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(prog="gap-reproduce") subparsers = parser.add_subparsers(dest="command", required=True) validate = subparsers.add_parser("validate-data") validate.add_argument("dataset", type=Path) validate.add_argument("--output", type=Path) summarize = subparsers.add_parser("summarize-run") summarize.add_argument("run_dir", type=Path) summarize.add_argument("--output", type=Path) align = subparsers.add_parser("align-release") align.add_argument("run_dir", type=Path) align.add_argument("dataset", type=Path) align.add_argument("--output", type=Path) generate = subparsers.add_parser("generate-kernel") generate.add_argument("--dataset", type=Path, required=True) generate.add_argument("--item-id", required=True) generate.add_argument("--run-dir", type=Path, required=True) generate.add_argument("--proposer-model", default="o3") generate.add_argument("--judge-model", default="o3") generate.add_argument( "--seed", type=int, help="optional provider seed; omitted by default for o3 compatibility", ) surfaces = subparsers.add_parser("generate-surfaces") surfaces.add_argument("--dataset", type=Path, required=True) surfaces.add_argument("--item-id", required=True) surfaces.add_argument("--run-dir", type=Path, required=True) surfaces.add_argument("--proposer-model", default="o3") export = subparsers.add_parser("export-release") export.add_argument("--source-dataset", type=Path, required=True) export.add_argument("--surface-run-dir", type=Path, required=True) export.add_argument("--kernel-run-dir", type=Path, required=True) export.add_argument("--output-root", type=Path, required=True) export.add_argument("--allow-partial", action="store_true") export.add_argument( "--item-id", action="append", dest="item_ids", help="export only this item ID; repeat to select multiple items", ) return parser def main() -> None: args = build_parser().parse_args() if args.command == "validate-data": write_or_print(validate_public_dataset(args.dataset), args.output) elif args.command == "summarize-run": write_or_print(summarize_run(args.run_dir), args.output) elif args.command == "align-release": write_or_print( align_run_to_release(args.run_dir, args.dataset), args.output, ) elif args.command == "generate-kernel": asyncio.run(generate_kernel(args)) elif args.command == "generate-surfaces": asyncio.run(generate_surfaces(args)) elif args.command == "export-release": manifest = export_release( source_dataset=args.source_dataset, surface_run_root=args.surface_run_dir, kernel_run_root=args.kernel_run_dir, output_root=args.output_root, allow_partial=args.allow_partial, item_ids=set(args.item_ids) if args.item_ids else None, ) print(json.dumps(manifest, ensure_ascii=False, indent=2)) else: raise AssertionError(args.command) if __name__ == "__main__": main()