summaryrefslogtreecommitdiff
path: root/src/gap_pipeline/paper_pipeline.py
diff options
context:
space:
mode:
Diffstat (limited to 'src/gap_pipeline/paper_pipeline.py')
-rw-r--r--src/gap_pipeline/paper_pipeline.py90
1 files changed, 59 insertions, 31 deletions
diff --git a/src/gap_pipeline/paper_pipeline.py b/src/gap_pipeline/paper_pipeline.py
index 84f34bf..c90f3f6 100644
--- a/src/gap_pipeline/paper_pipeline.py
+++ b/src/gap_pipeline/paper_pipeline.py
@@ -133,6 +133,7 @@ class PaperKernelPipeline:
methods: MethodPlan,
*,
version: int,
+ previous_replacements: ReplacementPlan | None = None,
feedback: str = "",
) -> ReplacementPlan:
request_id = f"{item.item_id}.stage3.replacement.v{version:02d}"
@@ -145,10 +146,13 @@ class PaperKernelPipeline:
item,
dag,
methods,
+ previous_replacements=previous_replacements,
feedback=feedback,
),
)
).validate_against(dag)
+ if previous_replacements is not None:
+ replacements.validate_repair_of(previous_replacements)
self.store.write_stage(
f"03_replacement_v{version:02d}",
replacements,
@@ -164,6 +168,8 @@ class PaperKernelPipeline:
replacements: ReplacementPlan,
*,
version: int,
+ previous_diffused: DiffusedProof | None = None,
+ feedback: str = "",
) -> DiffusedProof:
request_id = f"{item.item_id}.stage4.diffusion.v{version:02d}"
diffused = DiffusedProof.model_validate(
@@ -171,7 +177,14 @@ class PaperKernelPipeline:
self.proposer,
request_id=request_id,
system_prompt=DIFFUSION_SYSTEM,
- user_prompt=diffusion_user(item, dag, methods, replacements),
+ user_prompt=diffusion_user(
+ item,
+ dag,
+ methods,
+ replacements,
+ previous_diffused=previous_diffused,
+ feedback=feedback,
+ ),
)
).validate_against(dag, methods)
self.store.write_stage(
@@ -188,6 +201,8 @@ class PaperKernelPipeline:
diffused: DiffusedProof,
*,
version: int,
+ previous_variant: RenderedVariant | None = None,
+ feedback: str = "",
) -> RenderedVariant:
request_id = f"{self.store.item_id}.stage5.render.v{version:02d}"
variant = RenderedVariant.model_validate(
@@ -195,7 +210,12 @@ class PaperKernelPipeline:
self.proposer,
request_id=request_id,
system_prompt=RENDER_SYSTEM,
- user_prompt=render_user(replacements, diffused),
+ user_prompt=render_user(
+ replacements,
+ diffused,
+ previous_variant=previous_variant,
+ feedback=feedback,
+ ),
)
).validate_against(dag, diffused)
self.store.write_stage(
@@ -212,6 +232,7 @@ class PaperKernelPipeline:
methods: MethodPlan,
*,
version: int,
+ previous_bundle: CandidateBundle | None = None,
feedback: str = "",
) -> CandidateBundle:
replacements = await self.generate_replacements(
@@ -219,6 +240,11 @@ class PaperKernelPipeline:
dag,
methods,
version=version,
+ previous_replacements=(
+ previous_bundle.replacement_plan
+ if previous_bundle is not None
+ else None
+ ),
feedback=feedback,
)
diffused = await self.diffuse_dag(
@@ -227,18 +253,34 @@ class PaperKernelPipeline:
methods,
replacements,
version=version,
+ previous_diffused=(
+ previous_bundle.diffused_proof
+ if previous_bundle is not None
+ else None
+ ),
+ feedback=feedback,
)
variant = await self.render_variant(
dag,
replacements,
diffused,
version=version,
+ previous_variant=(
+ previous_bundle.variant if previous_bundle is not None else None
+ ),
+ feedback=feedback,
)
- return CandidateBundle(
+ bundle = CandidateBundle(
replacement_plan=replacements,
diffused_proof=diffused,
variant=variant,
)
+ if (
+ previous_bundle is not None
+ and sha256_payload(bundle) == sha256_payload(previous_bundle)
+ ):
+ raise ValueError("repair pass returned an unchanged candidate bundle")
+ return bundle
async def _judge_once(
self,
@@ -251,36 +293,21 @@ class PaperKernelPipeline:
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,
- ),
- )
+ 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,
+ ),
)
- 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}"
)
+ return verdict.validate_coverage(dag)
@staticmethod
def _feedback(verdicts: list[JudgeVerdict]) -> str:
@@ -369,6 +396,7 @@ class PaperKernelPipeline:
dag,
methods,
version=iteration + 1,
+ previous_bundle=bundle,
feedback=self._feedback(verdicts),
)
repaired_from_previous = True