summaryrefslogtreecommitdiff
path: root/src/gap_pipeline/pipeline.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/pipeline.py
Add minimal GAP reproduction package
Diffstat (limited to 'src/gap_pipeline/pipeline.py')
-rw-r--r--src/gap_pipeline/pipeline.py459
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,
+ )