summaryrefslogtreecommitdiff
path: root/src/gap_pipeline/paper_pipeline.py
diff options
context:
space:
mode:
authorOscar Wan <oscarwan@stanford.edu>2026-07-24 20:45:42 -0700
committerOscar Wan <oscarwan@stanford.edu>2026-07-24 20:45:42 -0700
commit15efc30e9e7179accd30375d3edb2e34a3b4dc5f (patch)
treeba1c49eb128e7906ba451141723d763e1dcda53a /src/gap_pipeline/paper_pipeline.py
parent708f2af9c6985e9cb5cd53e434a7d3b8dfa2b4ac (diff)
updated generation process
Diffstat (limited to 'src/gap_pipeline/paper_pipeline.py')
-rw-r--r--src/gap_pipeline/paper_pipeline.py395
1 files changed, 395 insertions, 0 deletions
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)