diff options
Diffstat (limited to 'src/gap_pipeline/cli.py')
| -rw-r--r-- | src/gap_pipeline/cli.py | 168 |
1 files changed, 168 insertions, 0 deletions
diff --git a/src/gap_pipeline/cli.py b/src/gap_pipeline/cli.py new file mode 100644 index 0000000..105cc97 --- /dev/null +++ b/src/gap_pipeline/cli.py @@ -0,0 +1,168 @@ +"""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() |
