summaryrefslogtreecommitdiff
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/gap_pipeline/clients.py17
-rw-r--r--src/gap_pipeline/paper_pipeline.py212
2 files changed, 150 insertions, 79 deletions
diff --git a/src/gap_pipeline/clients.py b/src/gap_pipeline/clients.py
index 3b3e40f..bc35742 100644
--- a/src/gap_pipeline/clients.py
+++ b/src/gap_pipeline/clients.py
@@ -8,6 +8,8 @@ from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Protocol
+from pydantic import BaseModel
+
@dataclass(frozen=True)
class ModelResponse:
@@ -27,6 +29,7 @@ class JsonLLM(Protocol):
system_prompt: str,
user_prompt: str,
request_id: str,
+ response_model: type[BaseModel] | None = None,
) -> ModelResponse: ...
@@ -61,14 +64,24 @@ class OpenAIJsonClient:
system_prompt: str,
user_prompt: str,
request_id: str,
+ response_model: type[BaseModel] | None = None,
) -> ModelResponse:
+ if response_model is None:
+ response_format: dict[str, Any] = {"type": "json_object"}
+ else:
+ from openai.lib._parsing._completions import (
+ type_to_response_format_param,
+ )
+
+ response_format = type_to_response_format_param(response_model)
+
kwargs: dict[str, Any] = {
"model": self.model,
"messages": [
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_prompt},
],
- "response_format": {"type": "json_object"},
+ "response_format": response_format,
"max_completion_tokens": self.max_completion_tokens,
}
if self.seed is not None:
@@ -111,6 +124,7 @@ class ReplayClient:
system_prompt: str,
user_prompt: str,
request_id: str,
+ response_model: type[BaseModel] | None = None,
) -> ModelResponse:
if request_id not in self._responses:
raise KeyError(f"no replay response for {request_id}")
@@ -145,6 +159,7 @@ class ScriptedClient:
system_prompt: str,
user_prompt: str,
request_id: str,
+ response_model: type[BaseModel] | None = None,
) -> ModelResponse:
self.calls.append(request_id)
prefixes = sorted(
diff --git a/src/gap_pipeline/paper_pipeline.py b/src/gap_pipeline/paper_pipeline.py
index c90f3f6..ffe9f15 100644
--- a/src/gap_pipeline/paper_pipeline.py
+++ b/src/gap_pipeline/paper_pipeline.py
@@ -3,8 +3,10 @@
from __future__ import annotations
import asyncio
+from collections.abc import Callable
+from typing import TypeVar
-from pydantic import BaseModel, ConfigDict, model_validator
+from pydantic import BaseModel, ConfigDict, ValidationError, model_validator
from .clients import JsonLLM
from .kernel_models import (
@@ -36,6 +38,9 @@ from .models import CanonicalItem, ModelCallRecord
from .store import RunStore, sha256_payload
+ValidatedModel = TypeVar("ValidatedModel", bound=BaseModel)
+
+
class PaperPipelineConfig(BaseModel):
model_config = ConfigDict(extra="forbid")
@@ -80,11 +85,13 @@ class PaperKernelPipeline:
request_id: str,
system_prompt: str,
user_prompt: str,
+ response_model: type[BaseModel] | None = None,
) -> dict:
response = await client.generate_json(
system_prompt=system_prompt,
user_prompt=user_prompt,
request_id=request_id,
+ response_model=response_model,
)
self.store.write_call(
ModelCallRecord(
@@ -100,30 +107,75 @@ class PaperKernelPipeline:
)
return response.data
+ async def _call_validated(
+ self,
+ client: JsonLLM,
+ *,
+ request_id: str,
+ system_prompt: str,
+ user_prompt: str,
+ response_model: type[ValidatedModel],
+ validate: Callable[[ValidatedModel], ValidatedModel] | None = None,
+ ) -> tuple[ValidatedModel, str]:
+ """Request a typed response, retrying the identical prompt once if invalid."""
+
+ last_error: ValidationError | ValueError | None = None
+ for attempt in range(2):
+ attempt_request_id = (
+ request_id
+ if attempt == 0
+ else f"{request_id}.validation_retry{attempt:02d}"
+ )
+ data = await self._call(
+ client,
+ request_id=attempt_request_id,
+ system_prompt=system_prompt,
+ user_prompt=user_prompt,
+ response_model=response_model,
+ )
+ try:
+ result = response_model.model_validate(data)
+ if validate is not None:
+ result = validate(result)
+ return result, attempt_request_id
+ except (ValidationError, ValueError) as exc:
+ last_error = exc
+
+ raise ValueError(
+ f"{request_id} failed typed validation after two identical-prompt attempts"
+ ) from last_error
+
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),
- )
+ dag, accepted_request_id = await self._call_validated(
+ self.proposer,
+ request_id=request_id,
+ system_prompt=DAG_SYSTEM,
+ user_prompt=dag_user(item),
+ response_model=ProofDAG,
+ )
+ self.store.write_stage(
+ "01_proof_dag",
+ dag,
+ request_id=accepted_request_id,
)
- 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)
+ methods, accepted_request_id = await self._call_validated(
+ self.proposer,
+ request_id=request_id,
+ system_prompt=METHOD_SYSTEM,
+ user_prompt=method_user(dag),
+ response_model=MethodPlan,
+ validate=lambda value: value.validate_against(dag),
+ )
+ self.store.write_stage(
+ "02_method_plan",
+ methods,
+ request_id=accepted_request_id,
+ )
return methods
async def generate_replacements(
@@ -137,26 +189,30 @@ class PaperKernelPipeline:
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,
- previous_replacements=previous_replacements,
- feedback=feedback,
- ),
- )
- ).validate_against(dag)
- if previous_replacements is not None:
- replacements.validate_repair_of(previous_replacements)
+ def validate_replacements(value: ReplacementPlan) -> ReplacementPlan:
+ value.validate_against(dag)
+ if previous_replacements is not None:
+ value.validate_repair_of(previous_replacements)
+ return value
+
+ replacements, accepted_request_id = await self._call_validated(
+ self.proposer,
+ request_id=request_id,
+ system_prompt=REPLACEMENT_SYSTEM,
+ user_prompt=replacement_user(
+ item,
+ dag,
+ methods,
+ previous_replacements=previous_replacements,
+ feedback=feedback,
+ ),
+ response_model=ReplacementPlan,
+ validate=validate_replacements,
+ )
self.store.write_stage(
f"03_replacement_v{version:02d}",
replacements,
- request_id=request_id,
+ request_id=accepted_request_id,
)
return replacements
@@ -172,25 +228,25 @@ class PaperKernelPipeline:
feedback: str = "",
) -> 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,
- previous_diffused=previous_diffused,
- feedback=feedback,
- ),
- )
- ).validate_against(dag, methods)
+ diffused, accepted_request_id = await self._call_validated(
+ self.proposer,
+ request_id=request_id,
+ system_prompt=DIFFUSION_SYSTEM,
+ user_prompt=diffusion_user(
+ item,
+ dag,
+ methods,
+ replacements,
+ previous_diffused=previous_diffused,
+ feedback=feedback,
+ ),
+ response_model=DiffusedProof,
+ validate=lambda value: value.validate_against(dag, methods),
+ )
self.store.write_stage(
f"04_diffused_proof_v{version:02d}",
diffused,
- request_id=request_id,
+ request_id=accepted_request_id,
)
return diffused
@@ -205,23 +261,23 @@ class PaperKernelPipeline:
feedback: str = "",
) -> 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,
- previous_variant=previous_variant,
- feedback=feedback,
- ),
- )
- ).validate_against(dag, diffused)
+ variant, accepted_request_id = await self._call_validated(
+ self.proposer,
+ request_id=request_id,
+ system_prompt=RENDER_SYSTEM,
+ user_prompt=render_user(
+ replacements,
+ diffused,
+ previous_variant=previous_variant,
+ feedback=feedback,
+ ),
+ response_model=RenderedVariant,
+ validate=lambda value: value.validate_against(dag, diffused),
+ )
self.store.write_stage(
f"05_rendered_variant_v{version:02d}",
variant,
- request_id=request_id,
+ request_id=accepted_request_id,
)
return variant
@@ -294,20 +350,20 @@ class PaperKernelPipeline:
judge_id: int,
) -> JudgeVerdict:
request_id = f"{item.item_id}.verify.t{iteration:02d}.j{judge_id}"
- verdict = JudgeVerdict.model_validate(
- await self._call(
- judge,
- request_id=request_id,
- system_prompt=JUDGE_SYSTEM,
- user_prompt=judge_user(
- item,
- methods,
- bundle.replacement_plan,
- bundle.variant,
- ),
- )
+ verdict, _ = await self._call_validated(
+ judge,
+ request_id=request_id,
+ system_prompt=JUDGE_SYSTEM,
+ user_prompt=judge_user(
+ item,
+ methods,
+ bundle.replacement_plan,
+ bundle.variant,
+ ),
+ response_model=JudgeVerdict,
+ validate=lambda value: value.validate_coverage(dag),
)
- return verdict.validate_coverage(dag)
+ return verdict
@staticmethod
def _feedback(verdicts: list[JudgeVerdict]) -> str: