From 15efc30e9e7179accd30375d3edb2e34a3b4dc5f Mon Sep 17 00:00:00 2001 From: Oscar Wan Date: Fri, 24 Jul 2026 20:45:42 -0700 Subject: updated generation process --- src/gap_pipeline/paper_pipeline.py | 395 +++++++++++++++++++++++++++++++++++++ 1 file changed, 395 insertions(+) create mode 100644 src/gap_pipeline/paper_pipeline.py (limited to 'src/gap_pipeline/paper_pipeline.py') diff --git a/src/gap_pipeline/paper_pipeline.py b/src/gap_pipeline/paper_pipeline.py new file mode 100644 index 0000000..84f34bf --- /dev/null +++ b/src/gap_pipeline/paper_pipeline.py @@ -0,0 +1,395 @@ +"""Literal five-stage GAP kernel pipeline matching the manuscript operations.""" + +from __future__ import annotations + +import asyncio + +from pydantic import BaseModel, ConfigDict, model_validator + +from .clients import JsonLLM +from .kernel_models import ( + CandidateBundle, + DiffusedProof, + JudgeVerdict, + KernelRunResult, + MethodPlan, + ProofDAG, + RenderedVariant, + ReplacementPlan, + VerificationIteration, +) +from .kernel_prompts import ( + DAG_SYSTEM, + DIFFUSION_SYSTEM, + JUDGE_SYSTEM, + METHOD_SYSTEM, + RENDER_SYSTEM, + REPLACEMENT_SYSTEM, + dag_user, + diffusion_user, + judge_user, + method_user, + render_user, + replacement_user, +) +from .models import CanonicalItem, ModelCallRecord +from .store import RunStore, sha256_payload + + +class PaperPipelineConfig(BaseModel): + model_config = ConfigDict(extra="forbid") + + protocol_name: str = "gap-literal-five-stage-J5-K2-T15" + proposer_model: str + judge_model: str + judge_count: int = 5 + streak_length: int = 2 + max_iterations: int = 15 + + @model_validator(mode="after") + def enforce_protocol(self) -> "PaperPipelineConfig": + 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") + return self + + +class PaperKernelPipeline: + """Execute all five paper stages as typed, separately auditable calls.""" + + def __init__( + self, + *, + proposer: JsonLLM, + judges: list[JsonLLM], + store: RunStore, + config: PaperPipelineConfig, + ) -> None: + if len(judges) != config.judge_count: + raise ValueError( + f"expected {config.judge_count} judges, 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, + ) -> dict: + 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.data + + async def construct_dag(self, item: CanonicalItem) -> ProofDAG: + request_id = f"{item.item_id}.stage1.dag" + dag = ProofDAG.model_validate( + await self._call( + self.proposer, + request_id=request_id, + system_prompt=DAG_SYSTEM, + user_prompt=dag_user(item), + ) + ) + self.store.write_stage("01_proof_dag", dag, request_id=request_id) + return dag + + async def summarize_methods(self, dag: ProofDAG) -> MethodPlan: + request_id = f"{self.store.item_id}.stage2.methods" + methods = MethodPlan.model_validate( + await self._call( + self.proposer, + request_id=request_id, + system_prompt=METHOD_SYSTEM, + user_prompt=method_user(dag), + ) + ).validate_against(dag) + self.store.write_stage("02_method_plan", methods, request_id=request_id) + return methods + + async def generate_replacements( + self, + item: CanonicalItem, + dag: ProofDAG, + methods: MethodPlan, + *, + version: int, + feedback: str = "", + ) -> ReplacementPlan: + request_id = f"{item.item_id}.stage3.replacement.v{version:02d}" + replacements = ReplacementPlan.model_validate( + await self._call( + self.proposer, + request_id=request_id, + system_prompt=REPLACEMENT_SYSTEM, + user_prompt=replacement_user( + item, + dag, + methods, + feedback=feedback, + ), + ) + ).validate_against(dag) + self.store.write_stage( + f"03_replacement_v{version:02d}", + replacements, + request_id=request_id, + ) + return replacements + + async def diffuse_dag( + self, + item: CanonicalItem, + dag: ProofDAG, + methods: MethodPlan, + replacements: ReplacementPlan, + *, + version: int, + ) -> DiffusedProof: + request_id = f"{item.item_id}.stage4.diffusion.v{version:02d}" + diffused = DiffusedProof.model_validate( + await self._call( + self.proposer, + request_id=request_id, + system_prompt=DIFFUSION_SYSTEM, + user_prompt=diffusion_user(item, dag, methods, replacements), + ) + ).validate_against(dag, methods) + self.store.write_stage( + f"04_diffused_proof_v{version:02d}", + diffused, + request_id=request_id, + ) + return diffused + + async def render_variant( + self, + dag: ProofDAG, + replacements: ReplacementPlan, + diffused: DiffusedProof, + *, + version: int, + ) -> RenderedVariant: + request_id = f"{self.store.item_id}.stage5.render.v{version:02d}" + variant = RenderedVariant.model_validate( + await self._call( + self.proposer, + request_id=request_id, + system_prompt=RENDER_SYSTEM, + user_prompt=render_user(replacements, diffused), + ) + ).validate_against(dag, diffused) + self.store.write_stage( + f"05_rendered_variant_v{version:02d}", + variant, + request_id=request_id, + ) + return variant + + async def build_bundle( + self, + item: CanonicalItem, + dag: ProofDAG, + methods: MethodPlan, + *, + version: int, + feedback: str = "", + ) -> CandidateBundle: + replacements = await self.generate_replacements( + item, + dag, + methods, + version=version, + feedback=feedback, + ) + diffused = await self.diffuse_dag( + item, + dag, + methods, + replacements, + version=version, + ) + variant = await self.render_variant( + dag, + replacements, + diffused, + version=version, + ) + return CandidateBundle( + replacement_plan=replacements, + diffused_proof=diffused, + variant=variant, + ) + + async def _judge_once( + self, + judge: JsonLLM, + *, + item: CanonicalItem, + dag: ProofDAG, + methods: MethodPlan, + bundle: CandidateBundle, + iteration: int, + judge_id: int, + ) -> JudgeVerdict: + format_feedback = "" + for attempt in range(1, 4): + request_id = ( + f"{item.item_id}.verify.t{iteration:02d}." + f"j{judge_id}.a{attempt}" + ) + verdict = JudgeVerdict.model_validate( + await self._call( + judge, + request_id=request_id, + system_prompt=JUDGE_SYSTEM, + user_prompt=judge_user( + item, + dag, + methods, + bundle.replacement_plan, + bundle.diffused_proof, + bundle.variant, + format_feedback=format_feedback, + ), + ) + ) + try: + return verdict.validate_coverage(dag, bundle.replacement_plan) + except ValueError as exc: + format_feedback = str(exc) + raise ValueError( + f"judge {judge_id} failed coverage after three format attempts: " + f"{format_feedback}" + ) + + @staticmethod + def _feedback(verdicts: list[JudgeVerdict]) -> str: + rows = [] + for index, verdict in enumerate(verdicts, start=1): + if verdict.verdict == "reject": + rows.append( + f"judge {index}: {verdict.blocking_issues}; " + f"suggested repair: {verdict.patch_suggestion}" + ) + return "\n".join(rows) + + async def verify( + self, + item: CanonicalItem, + dag: ProofDAG, + methods: MethodPlan, + initial_bundle: CandidateBundle, + ) -> KernelRunResult: + bundle = initial_bundle + iterations: list[VerificationIteration] = [] + pass_streak = 0 + streak_sha: str | None = None + repaired_from_previous = False + + for iteration in range(1, self.config.max_iterations + 1): + bundle_sha = sha256_payload(bundle) + candidate_sha = sha256_payload(bundle.variant) + verdicts = list( + await asyncio.gather( + *( + self._judge_once( + judge, + item=item, + dag=dag, + methods=methods, + bundle=bundle, + 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_sha not in {None, bundle_sha}: + raise AssertionError("pass streak crossed provenance versions") + streak_sha = bundle_sha + pass_streak += 1 + else: + pass_streak = 0 + streak_sha = None + + record = VerificationIteration( + iteration=iteration, + bundle_sha256=bundle_sha, + candidate_sha256=candidate_sha, + verdicts=verdicts, + unanimous=unanimous, + pass_streak_after=pass_streak, + repaired_from_previous=repaired_from_previous, + ) + 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", + proof_dag=dag, + method_plan=methods, + accepted_replacement_plan=bundle.replacement_plan, + accepted_diffused_proof=bundle.diffused_proof, + accepted_candidate=bundle.variant, + iterations=iterations, + accepted_candidate_sha256=candidate_sha, + accepted_bundle_sha256=bundle_sha, + ) + self.store.write_final(result) + return result + + if not unanimous and iteration < self.config.max_iterations: + bundle = await self.build_bundle( + item, + dag, + methods, + version=iteration + 1, + feedback=self._feedback(verdicts), + ) + repaired_from_previous = True + else: + repaired_from_previous = False + + result = KernelRunResult( + item_id=item.item_id, + status="rejected", + proof_dag=dag, + method_plan=methods, + iterations=iterations, + rejection_reason="no two consecutive unanimous rounds within T=15", + ) + self.store.write_final(result) + return result + + async def run(self, item: CanonicalItem) -> KernelRunResult: + self.store.write_input(item) + self.store.write_config(self.config) + dag = await self.construct_dag(item) + methods = await self.summarize_methods(dag) + bundle = await self.build_bundle(item, dag, methods, version=1) + return await self.verify(item, dag, methods, bundle) -- cgit v1.2.3