summaryrefslogtreecommitdiff
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/gap_pipeline/__init__.py4
-rw-r--r--src/gap_pipeline/e2e.py2
-rw-r--r--src/gap_pipeline/models.py126
-rw-r--r--src/gap_pipeline/pipeline.py36
-rw-r--r--src/gap_pipeline/prompts.py5
5 files changed, 152 insertions, 21 deletions
diff --git a/src/gap_pipeline/__init__.py b/src/gap_pipeline/__init__.py
index fe2e04c..fbc2adb 100644
--- a/src/gap_pipeline/__init__.py
+++ b/src/gap_pipeline/__init__.py
@@ -5,6 +5,8 @@ from .models import (
KernelCandidate,
KernelPlan,
KernelRunResult,
+ ProofPlanDAG,
+ ProofPlanNode,
)
from .pipeline import KernelPipeline, PipelineConfig
@@ -15,6 +17,8 @@ __all__ = [
"KernelPlan",
"KernelRunResult",
"PipelineConfig",
+ "ProofPlanDAG",
+ "ProofPlanNode",
]
__version__ = "0.2.0"
diff --git a/src/gap_pipeline/e2e.py b/src/gap_pipeline/e2e.py
index 971da3e..ded1381 100644
--- a/src/gap_pipeline/e2e.py
+++ b/src/gap_pipeline/e2e.py
@@ -25,7 +25,7 @@ def _ensure_fresh(path: Path) -> None:
def _review_accept() -> dict[str, str]:
return {
"verdict": "accept",
- "step_by_step_check": "both method labels are instantiated",
+ "step_by_step_check": "n1 passes; n2 passes",
"blocking_issues": "",
"patch_suggestion": "",
}
diff --git a/src/gap_pipeline/models.py b/src/gap_pipeline/models.py
index bf80c5b..6fab32f 100644
--- a/src/gap_pipeline/models.py
+++ b/src/gap_pipeline/models.py
@@ -2,6 +2,7 @@
from __future__ import annotations
+import re
from datetime import datetime, timezone
from typing import Any, Literal
@@ -67,6 +68,115 @@ class KernelPlan(StrictModel):
return self
+class ProofPlanNode(StrictModel):
+ node_id: str
+ method_label: str
+ dependencies: list[str] = Field(default_factory=list)
+
+ @model_validator(mode="after")
+ def validate_content(self) -> "ProofPlanNode":
+ if not self.node_id.strip():
+ raise ValueError("proof-plan node ID must be non-empty")
+ if not self.method_label.strip():
+ raise ValueError("proof-plan method label must be non-empty")
+ return self
+
+
+class ProofPlanDAG(StrictModel):
+ """Typed path-DAG induced by Prompt-A's ordered proof-plan steps."""
+
+ nodes: list[ProofPlanNode] = Field(min_length=1)
+ terminal_node_id: str
+ structure: Literal["path"] = "path"
+
+ @classmethod
+ def from_plan(cls, plan: KernelPlan) -> "ProofPlanDAG":
+ nodes = [
+ ProofPlanNode(
+ node_id=f"n{index}",
+ method_label=step,
+ dependencies=[] if index == 1 else [f"n{index - 1}"],
+ )
+ for index, step in enumerate(plan.core_steps, start=1)
+ ]
+ return cls(nodes=nodes, terminal_node_id=nodes[-1].node_id)
+
+ def method_labels_payload(self) -> list[str]:
+ """Keep the judge input a sequence while making node coverage auditable."""
+
+ return [f"{node.node_id}: {node.method_label}" for node in self.nodes]
+
+ def validate_plan(self, plan: KernelPlan) -> "ProofPlanDAG":
+ labels = [node.method_label for node in self.nodes]
+ if labels != plan.core_steps:
+ raise ValueError(
+ "proof-plan DAG labels must equal Prompt-A core_steps in order"
+ )
+ return self
+
+ @model_validator(mode="after")
+ def validate_graph(self) -> "ProofPlanDAG":
+ node_ids = [node.node_id for node in self.nodes]
+ if len(node_ids) != len(set(node_ids)):
+ raise ValueError("proof-plan DAG node IDs must be unique")
+
+ known = set(node_ids)
+ if self.terminal_node_id not in known:
+ raise ValueError("proof-plan DAG terminal node is unknown")
+
+ dependents: dict[str, list[str]] = {node_id: [] for node_id in node_ids}
+ indegree = {node.node_id: len(node.dependencies) for node in self.nodes}
+ for node in self.nodes:
+ if len(node.dependencies) != len(set(node.dependencies)):
+ raise ValueError(f"{node.node_id} has duplicate dependencies")
+ for dependency in node.dependencies:
+ if dependency not in known:
+ raise ValueError(
+ f"{node.node_id} depends on unknown node {dependency}"
+ )
+ if dependency == node.node_id:
+ raise ValueError(f"{node.node_id} cannot depend on itself")
+ dependents[dependency].append(node.node_id)
+
+ ready = [node_id for node_id in node_ids if indegree[node_id] == 0]
+ visited = 0
+ while ready:
+ node_id = ready.pop()
+ visited += 1
+ for dependent in dependents[node_id]:
+ indegree[dependent] -= 1
+ if indegree[dependent] == 0:
+ ready.append(dependent)
+ if visited != len(node_ids):
+ raise ValueError("proof-plan graph contains a cycle")
+
+ by_id = {node.node_id: node for node in self.nodes}
+ ancestors: set[str] = set()
+ pending = [self.terminal_node_id]
+ while pending:
+ node_id = pending.pop()
+ if node_id in ancestors:
+ continue
+ ancestors.add(node_id)
+ pending.extend(by_id[node_id].dependencies)
+ if ancestors != known:
+ raise ValueError("every proof-plan node must lead to the terminal node")
+
+ expected_ids = [f"n{index}" for index in range(1, len(self.nodes) + 1)]
+ if node_ids != expected_ids:
+ raise ValueError("path-DAG nodes must be ordered n1 through nN")
+ for index, node in enumerate(self.nodes):
+ expected_dependencies = [] if index == 0 else [node_ids[index - 1]]
+ if node.dependencies != expected_dependencies:
+ raise ValueError(
+ f"{node.node_id} must depend exactly on "
+ f"{expected_dependencies or 'no prior node'}"
+ )
+ if self.terminal_node_id != node_ids[-1]:
+ raise ValueError("path-DAG terminal must be the final ordered node")
+ return self
+
+
class KernelCandidate(StrictModel):
question: str
solution: str
@@ -92,6 +202,20 @@ class JudgeVerdict(StrictModel):
raise ValueError("reject verdict must identify a blocking issue")
return self
+ def validate_coverage(self, dag: ProofPlanDAG) -> "JudgeVerdict":
+ missing = [
+ node.node_id
+ for node in dag.nodes
+ if re.search(
+ rf"(?<![A-Za-z0-9_]){re.escape(node.node_id)}(?![A-Za-z0-9_])",
+ self.step_by_step_check,
+ )
+ is None
+ ]
+ if missing:
+ raise ValueError(f"judge check does not cover DAG nodes: {missing}")
+ return self
+
class IterationRecord(StrictModel):
iteration: int
@@ -105,6 +229,7 @@ class KernelRunResult(StrictModel):
item_id: str
status: Literal["accepted", "rejected"]
plan: KernelPlan
+ proof_dag: ProofPlanDAG
accepted_candidate: KernelCandidate | None = None
iterations: list[IterationRecord]
rejection_reason: str = ""
@@ -113,6 +238,7 @@ class KernelRunResult(StrictModel):
@model_validator(mode="after")
def validate_terminal_state(self) -> "KernelRunResult":
+ self.proof_dag.validate_plan(self.plan)
if self.status == "accepted" and self.accepted_candidate is None:
raise ValueError("accepted run must include its candidate")
if self.status == "rejected" and not self.rejection_reason:
diff --git a/src/gap_pipeline/pipeline.py b/src/gap_pipeline/pipeline.py
index 307ba81..76ae722 100644
--- a/src/gap_pipeline/pipeline.py
+++ b/src/gap_pipeline/pipeline.py
@@ -15,6 +15,7 @@ from .models import (
KernelPlan,
KernelRunResult,
ModelCallRecord,
+ ProofPlanDAG,
)
from .prompts import (
FIX_SYSTEM_PROMPT,
@@ -91,7 +92,10 @@ class KernelPipeline:
)
return response.data
- async def extract_plan(self, item: CanonicalItem) -> KernelPlan:
+ async def extract_plan(
+ self,
+ item: CanonicalItem,
+ ) -> tuple[KernelPlan, ProofPlanDAG]:
request_id = f"{item.item_id}.plan"
payload = await self._call(
self.proposer,
@@ -100,21 +104,10 @@ class KernelPipeline:
user_prompt=kernel_plan_user(item),
)
plan = KernelPlan.model_validate(payload)
- nodes = [
- {
- "node_id": f"n{index}",
- "method_label": step,
- "dependencies": [] if index == 1 else [f"n{index - 1}"],
- }
- for index, step in enumerate(plan.core_steps, start=1)
- ]
+ dag = ProofPlanDAG.from_plan(plan)
self.store.write_stage(
"01_proof_dag",
- {
- "nodes": nodes,
- "terminal_node_id": nodes[-1]["node_id"],
- "construction": "ordered core_steps returned verbatim by Prompt-A",
- },
+ dag,
request_id=request_id,
)
self.store.write_stage(
@@ -130,7 +123,7 @@ class KernelPipeline:
},
request_id=request_id,
)
- return plan
+ return plan, dag
async def generate_candidate(
self,
@@ -163,6 +156,7 @@ class KernelPipeline:
*,
item: CanonicalItem,
plan: KernelPlan,
+ dag: ProofPlanDAG,
candidate: KernelCandidate,
iteration: int,
judge_id: int,
@@ -172,9 +166,9 @@ class KernelPipeline:
judge,
request_id=request_id,
system_prompt=JUDGE_SYSTEM_PROMPT,
- user_prompt=judge_user(item, plan, candidate),
+ user_prompt=judge_user(item, plan, dag, candidate),
)
- return JudgeVerdict.model_validate(payload)
+ return JudgeVerdict.model_validate(payload).validate_coverage(dag)
async def _repair(
self,
@@ -220,8 +214,10 @@ class KernelPipeline:
*,
item: CanonicalItem,
plan: KernelPlan,
+ dag: ProofPlanDAG,
candidate: KernelCandidate,
) -> KernelRunResult:
+ dag.validate_plan(plan)
current = candidate
iterations: list[IterationRecord] = []
pass_streak = 0
@@ -236,6 +232,7 @@ class KernelPipeline:
judge,
item=item,
plan=plan,
+ dag=dag,
candidate=current,
iteration=iteration,
judge_id=judge_id,
@@ -269,6 +266,7 @@ class KernelPipeline:
item_id=item.item_id,
status="accepted",
plan=plan,
+ proof_dag=dag,
accepted_candidate=current,
iterations=iterations,
accepted_candidate_sha256=current_sha,
@@ -288,6 +286,7 @@ class KernelPipeline:
item_id=item.item_id,
status="rejected",
plan=plan,
+ proof_dag=dag,
iterations=iterations,
rejection_reason="no two consecutive unanimous rounds within T=15",
)
@@ -297,10 +296,11 @@ class KernelPipeline:
async def run(self, item: CanonicalItem) -> KernelRunResult:
self.store.write_input(item)
self.store.write_config(self.config)
- plan = await self.extract_plan(item)
+ plan, dag = await self.extract_plan(item)
candidate = await self.generate_candidate(item, plan)
return await self.verify_candidate(
item=item,
plan=plan,
+ dag=dag,
candidate=candidate,
)
diff --git a/src/gap_pipeline/prompts.py b/src/gap_pipeline/prompts.py
index 2ba87cf..30d65b2 100644
--- a/src/gap_pipeline/prompts.py
+++ b/src/gap_pipeline/prompts.py
@@ -8,7 +8,7 @@ from __future__ import annotations
import json
-from .models import CanonicalItem, KernelCandidate, KernelPlan
+from .models import CanonicalItem, KernelCandidate, KernelPlan, ProofPlanDAG
# Source: PutnamVariants@c3bed737370df2dbf73afd66bf6e86d4ece82d68
@@ -220,12 +220,13 @@ def kernel_generate_user(item: CanonicalItem, plan: KernelPlan) -> str:
def judge_user(
item: CanonicalItem,
plan: KernelPlan,
+ dag: ProofPlanDAG,
candidate: KernelCandidate,
) -> str:
return JUDGE_USER_TEMPLATE.format(
original_problem=item.problem,
original_solution=item.solution,
- method_labels=json.dumps(plan.core_steps, ensure_ascii=False),
+ method_labels=json.dumps(dag.method_labels_payload(), ensure_ascii=False),
slot_replacement=json.dumps(
{
key: value.model_dump(mode="json")