diff options
| author | Anonymous Authors <anonymous@invalid.example> | 2026-07-24 13:24:36 -0500 |
|---|---|---|
| committer | Anonymous Authors <anonymous@invalid.example> | 2026-07-24 13:24:36 -0500 |
| commit | db293f3606a97b3e417de27124858e134005acbd (patch) | |
| tree | 8efeedcd2033b82d1c90eb0cb84e134421ff1a8f /src/gap_pipeline/pipeline.py | |
Add minimal GAP reproduction package
Diffstat (limited to 'src/gap_pipeline/pipeline.py')
| -rw-r--r-- | src/gap_pipeline/pipeline.py | 459 |
1 files changed, 459 insertions, 0 deletions
diff --git a/src/gap_pipeline/pipeline.py b/src/gap_pipeline/pipeline.py new file mode 100644 index 0000000..2c0deb4 --- /dev/null +++ b/src/gap_pipeline/pipeline.py @@ -0,0 +1,459 @@ +"""Five-stage kernel generation and the J=5, K=2, T=15 verifier.""" + +from __future__ import annotations + +import asyncio +import hashlib +from pathlib import Path +from typing import Any + +from pydantic import BaseModel, ConfigDict, model_validator + +from .clients import JsonLLM +from .models import ( + CanonicalItem, + DiffusedProof, + IterationRecord, + JudgeVerdict, + KernelCandidate, + KernelRunResult, + MethodPlan, + ModelCallRecord, + ProofDAG, + ReplacementSpec, +) +from .prompts import ( + DAG_SYSTEM, + DIFFUSION_SYSTEM, + JUDGE_SYSTEM, + METHOD_SYSTEM, + QUESTION_SYSTEM, + REPAIR_SYSTEM, + REPLACEMENT_SYSTEM, + dag_user, + diffusion_user, + judge_user, + method_user, + question_user, + repair_user, + replacement_user, +) +from .store import RunStore, sha256_payload + + +class PipelineConfig(BaseModel): + model_config = ConfigDict(extra="forbid") + + protocol_name: str = "gap-J5-K2-T15" + proposer_model: str + judge_model: str + judge_count: int = 5 + streak_length: int = 2 + max_iterations: int = 15 + human_audit_fraction: float = 0.10 + audit_seed: int = 0 + difficulty_instruction: str = ( + "Preserve the proof plan and avoid trivializing the source problem. " + "Routine algebraic adaptation is allowed, but do not introduce a new " + "pivotal lemma or an independent proof strategy." + ) + + @model_validator(mode="after") + def enforce_submitted_protocol(self) -> "PipelineConfig": + if (self.judge_count, self.streak_length, self.max_iterations) != (5, 2, 15): + raise ValueError( + "the GAP protocol is fixed at J=5, K=2, T=15" + ) + if self.human_audit_fraction != 0.10: + raise ValueError("the GAP post-hoc audit fraction is 10%") + return self + + +class KernelPipeline: + def __init__( + self, + *, + proposer: JsonLLM, + judges: list[JsonLLM], + store: RunStore, + config: PipelineConfig, + ) -> None: + if len(judges) != config.judge_count: + raise ValueError( + f"expected {config.judge_count} judge clients, received {len(judges)}" + ) + self.proposer = proposer + self.judges = judges + self.store = store + self.config = config + + async def _call( + self, + client: JsonLLM, + *, + request_id: str, + system_prompt: str, + user_prompt: str, + response_type: type[BaseModel], + ) -> BaseModel: + response = await client.generate_json( + system_prompt=system_prompt, + user_prompt=user_prompt, + request_id=request_id, + ) + self.store.write_call( + ModelCallRecord( + request_id=request_id, + model=response.model, + system_prompt=system_prompt, + user_prompt=user_prompt, + response_data=response.data, + raw_text=response.raw_text, + provider_response_id=response.response_id, + usage=response.usage, + ) + ) + return response_type.model_validate(response.data) + + async def construct_dag(self, item: CanonicalItem) -> ProofDAG: + request_id = f"{item.item_id}.stage1.dag" + dag = await self._call( + self.proposer, + request_id=request_id, + system_prompt=DAG_SYSTEM, + user_prompt=dag_user(item), + response_type=ProofDAG, + ) + assert isinstance(dag, ProofDAG) + self.store.write_stage("01_proof_dag", dag, request_id=request_id) + return dag + + async def summarize_methods(self, item: CanonicalItem, dag: ProofDAG) -> MethodPlan: + request_id = f"{item.item_id}.stage2.methods" + plan = await self._call( + self.proposer, + request_id=request_id, + system_prompt=METHOD_SYSTEM, + user_prompt=method_user(dag), + response_type=MethodPlan, + ) + assert isinstance(plan, MethodPlan) + plan.validate_against(dag) + self.store.write_stage("02_method_plan", plan, request_id=request_id) + return plan + + async def generate_replacement( + self, + item: CanonicalItem, + dag: ProofDAG, + plan: MethodPlan, + ) -> ReplacementSpec: + request_id = f"{item.item_id}.stage3.replacement" + replacement = await self._call( + self.proposer, + request_id=request_id, + system_prompt=REPLACEMENT_SYSTEM, + user_prompt=replacement_user( + item, + dag, + plan, + difficulty_instruction=self.config.difficulty_instruction, + ), + response_type=ReplacementSpec, + ) + assert isinstance(replacement, ReplacementSpec) + replacement.validate_against(dag) + self.store.write_stage( + "03_replacement", + replacement, + request_id=request_id, + ) + return replacement + + async def diffuse_dag( + self, + item: CanonicalItem, + dag: ProofDAG, + plan: MethodPlan, + replacement: ReplacementSpec, + ) -> DiffusedProof: + request_id = f"{item.item_id}.stage4.diffusion" + diffused = await self._call( + self.proposer, + request_id=request_id, + system_prompt=DIFFUSION_SYSTEM, + user_prompt=diffusion_user(item, dag, plan, replacement), + response_type=DiffusedProof, + ) + assert isinstance(diffused, DiffusedProof) + diffused.validate_against(dag, plan) + self.store.write_stage( + "04_diffused_proof", + diffused, + request_id=request_id, + ) + return diffused + + async def render_question( + self, + item: CanonicalItem, + dag: ProofDAG, + plan: MethodPlan, + replacement: ReplacementSpec, + diffused: DiffusedProof, + ) -> KernelCandidate: + request_id = f"{item.item_id}.stage5.question" + candidate = await self._call( + self.proposer, + request_id=request_id, + system_prompt=QUESTION_SYSTEM, + user_prompt=question_user( + item, + dag, + plan, + replacement, + diffused.model_dump(mode="json"), + ), + response_type=KernelCandidate, + ) + assert isinstance(candidate, KernelCandidate) + candidate.validate_against(dag, plan) + self.store.write_stage( + "05_draft_candidate", + candidate, + request_id=request_id, + ) + return candidate + + async def build_candidate( + self, + item: CanonicalItem, + ) -> tuple[ProofDAG, MethodPlan, KernelCandidate]: + dag = await self.construct_dag(item) + plan = await self.summarize_methods(item, dag) + replacement = await self.generate_replacement(item, dag, plan) + diffused = await self.diffuse_dag(item, dag, plan, replacement) + candidate = await self.render_question( + item, + dag, + plan, + replacement, + diffused, + ) + return dag, plan, candidate + + async def _judge_once( + self, + judge: JsonLLM, + *, + item: CanonicalItem, + plan: MethodPlan, + candidate: KernelCandidate, + iteration: int, + judge_id: int, + ) -> JudgeVerdict: + request_id = f"{item.item_id}.verify.t{iteration:02d}.j{judge_id}" + verdict = await self._call( + judge, + request_id=request_id, + system_prompt=JUDGE_SYSTEM, + user_prompt=judge_user( + item, + plan, + candidate, + judge_id=judge_id, + iteration=iteration, + ), + response_type=JudgeVerdict, + ) + assert isinstance(verdict, JudgeVerdict) + return verdict + + async def _repair( + self, + *, + item: CanonicalItem, + dag: ProofDAG, + plan: MethodPlan, + candidate: KernelCandidate, + verdicts: list[JudgeVerdict], + iteration: int, + ) -> KernelCandidate: + request_id = f"{item.item_id}.repair.after_t{iteration:02d}" + issues = [ + { + "judge_id": judge_id, + "blocking_issues": verdict.blocking_issues, + "patch_suggestion": verdict.patch_suggestion, + } + for judge_id, verdict in enumerate(verdicts, start=1) + if verdict.verdict == "reject" + ] + repaired = await self._call( + self.proposer, + request_id=request_id, + system_prompt=REPAIR_SYSTEM, + user_prompt=repair_user(item, dag, plan, candidate, issues), + response_type=KernelCandidate, + ) + assert isinstance(repaired, KernelCandidate) + repaired.validate_against(dag, plan) + if sha256_payload(repaired) == sha256_payload(candidate): + raise ValueError("repair call returned a byte-equivalent candidate") + self.store.write_stage( + f"repair_after_{iteration:02d}", + repaired, + request_id=request_id, + ) + return repaired + + def _audit_selected(self, item_id: str) -> bool: + digest = hashlib.sha256( + f"{self.config.audit_seed}:{item_id}".encode("utf-8") + ).digest() + draw = int.from_bytes(digest[:8], "big") / float(2**64) + return draw < self.config.human_audit_fraction + + async def verify_candidate( + self, + *, + item: CanonicalItem, + dag: ProofDAG, + plan: MethodPlan, + candidate: KernelCandidate, + ) -> KernelRunResult: + iterations: list[IterationRecord] = [] + pass_streak = 0 + current = candidate + streak_candidate_sha: str | None = None + + for iteration in range(1, self.config.max_iterations + 1): + current_sha = sha256_payload(current) + verdicts = list( + await asyncio.gather( + *( + self._judge_once( + judge, + item=item, + plan=plan, + candidate=current, + iteration=iteration, + judge_id=judge_id, + ) + for judge_id, judge in enumerate(self.judges, start=1) + ) + ) + ) + unanimous = all(verdict.verdict == "accept" for verdict in verdicts) + + if unanimous: + if streak_candidate_sha not in {None, current_sha}: + raise AssertionError("pass streak crossed candidate versions") + streak_candidate_sha = current_sha + pass_streak += 1 + record = IterationRecord( + iteration=iteration, + candidate_sha256=current_sha, + verdicts=verdicts, + unanimous=True, + pass_streak_after=pass_streak, + ) + iterations.append(record) + self.store.write_iteration(iteration, record) + if pass_streak >= self.config.streak_length: + result = KernelRunResult( + item_id=item.item_id, + status="accepted", + accepted_candidate=current, + iterations=iterations, + ) + self.store.write_final( + { + **result.model_dump(mode="json"), + "accepted_candidate_sha256": current_sha, + "human_audit_selected": self._audit_selected(item.item_id), + } + ) + return result + continue + + pass_streak = 0 + streak_candidate_sha = None + if iteration == self.config.max_iterations: + record = IterationRecord( + iteration=iteration, + candidate_sha256=current_sha, + verdicts=verdicts, + unanimous=False, + pass_streak_after=0, + ) + iterations.append(record) + self.store.write_iteration(iteration, record) + break + + repaired = await self._repair( + item=item, + dag=dag, + plan=plan, + candidate=current, + verdicts=verdicts, + iteration=iteration, + ) + repaired_sha = sha256_payload(repaired) + record = IterationRecord( + iteration=iteration, + candidate_sha256=current_sha, + verdicts=verdicts, + unanimous=False, + pass_streak_after=0, + repaired=True, + repaired_candidate_sha256=repaired_sha, + ) + iterations.append(record) + self.store.write_iteration(iteration, record) + current = repaired + + result = KernelRunResult( + item_id=item.item_id, + status="rejected", + iterations=iterations, + rejection_reason=( + f"no two consecutive unanimous passes within " + f"{self.config.max_iterations} iterations" + ), + ) + self.store.write_final( + { + **result.model_dump(mode="json"), + "human_audit_selected": False, + } + ) + return result + + async def run(self, item: CanonicalItem) -> KernelRunResult: + self.store.write_input(item) + self.store.write_config(self.config) + dag, plan, candidate = await self.build_candidate(item) + return await self.verify_candidate( + item=item, + dag=dag, + plan=plan, + candidate=candidate, + ) + + +def make_pipeline( + *, + proposer: JsonLLM, + judge_factory: Any, + run_dir: Path, + item_id: str, + config: PipelineConfig, +) -> KernelPipeline: + judges = [judge_factory(judge_id) for judge_id in range(1, 6)] + return KernelPipeline( + proposer=proposer, + judges=judges, + store=RunStore(run_dir, item_id), + config=config, + ) |
