summaryrefslogtreecommitdiff
path: root/src/gap_pipeline/cli.py
diff options
context:
space:
mode:
authorAnonymous Authors <anonymous@invalid.example>2026-07-24 13:24:36 -0500
committerAnonymous Authors <anonymous@invalid.example>2026-07-24 13:24:36 -0500
commitdb293f3606a97b3e417de27124858e134005acbd (patch)
tree8efeedcd2033b82d1c90eb0cb84e134421ff1a8f /src/gap_pipeline/cli.py
Add minimal GAP reproduction package
Diffstat (limited to 'src/gap_pipeline/cli.py')
-rw-r--r--src/gap_pipeline/cli.py168
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()