Added structure refiner (codex)
This commit is contained in:
114
structure_refiner.py
Normal file
114
structure_refiner.py
Normal file
@@ -0,0 +1,114 @@
|
||||
import json
|
||||
import os
|
||||
import traceback
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, ValidationError, model_validator
|
||||
|
||||
from agent import Agent, AgentMessage
|
||||
from structure_builder import ImageElement, Structure, StructureElement
|
||||
from windowizer import Window, Windowizer
|
||||
|
||||
|
||||
class _StrictModel(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid", strict=True)
|
||||
|
||||
|
||||
class _RefinerContext(_StrictModel):
|
||||
current_section: str | None = Field(default=None, max_length=120)
|
||||
current_subsection: str | None = Field(default=None, max_length=120)
|
||||
recent_headings: list[str] = Field(default_factory=list, max_length=6)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_recent_headings(self) -> "_RefinerContext":
|
||||
if any(not heading.strip() or len(heading) > 120 for heading in self.recent_headings):
|
||||
raise ValueError("recent_headings must contain non-empty strings up to 120 characters")
|
||||
return self
|
||||
|
||||
|
||||
class _RefineResult(_StrictModel):
|
||||
elements: list[StructureElement]
|
||||
context: _RefinerContext
|
||||
|
||||
|
||||
class StructureRefiner:
|
||||
"""Perform a final editing pass over a document structure."""
|
||||
|
||||
DEBUG_ID = 0
|
||||
MAX_RETRIES = 5
|
||||
|
||||
def __init__(self, agent: Agent, windowizer: Windowizer[StructureElement]) -> None:
|
||||
self._agent = agent
|
||||
self._windowizer = windowizer
|
||||
with open("prompts/structure_refiner.md", "r", encoding="utf-8") as f:
|
||||
self._system_prompt = AgentMessage(content=f.read(), role="system")
|
||||
|
||||
@staticmethod
|
||||
def _validate_result(
|
||||
result: _RefineResult,
|
||||
window: Window[StructureElement],
|
||||
) -> None:
|
||||
input_image_ids = {
|
||||
element.event_id
|
||||
for element in window.present
|
||||
if isinstance(element, ImageElement)
|
||||
}
|
||||
for element in result.elements:
|
||||
if isinstance(element, ImageElement) and element.event_id not in input_image_ids:
|
||||
raise ValueError("image.event_id must come from present")
|
||||
|
||||
def _refine_window(self, window: Window[StructureElement]) -> _RefineResult:
|
||||
messages = [
|
||||
self._system_prompt,
|
||||
AgentMessage(
|
||||
content=json.dumps(
|
||||
window.model_dump(mode="json"),
|
||||
indent=2,
|
||||
ensure_ascii=False,
|
||||
),
|
||||
role="user",
|
||||
),
|
||||
]
|
||||
|
||||
debug_dir = None
|
||||
if os.path.isdir("debug"):
|
||||
debug_dir = f"debug/StructureRefiner/{StructureRefiner.DEBUG_ID}"
|
||||
StructureRefiner.DEBUG_ID += 1
|
||||
os.makedirs(debug_dir, exist_ok=True)
|
||||
with open(f"{debug_dir}/request.txt", "w", encoding="utf-8") as f:
|
||||
f.write(messages[1].content)
|
||||
|
||||
retries_left = self.MAX_RETRIES
|
||||
while retries_left > 0:
|
||||
retries_left -= 1
|
||||
response = self._agent.completion(
|
||||
messages=messages,
|
||||
response_format=_RefineResult,
|
||||
)
|
||||
if debug_dir:
|
||||
with open(
|
||||
f"{debug_dir}/{retries_left}-retries-left.txt",
|
||||
"w",
|
||||
encoding="utf-8",
|
||||
) as f:
|
||||
f.write(response)
|
||||
|
||||
try:
|
||||
result = _RefineResult.model_validate_json(response)
|
||||
self._validate_result(result, window)
|
||||
return result
|
||||
except (ValidationError, ValueError, TypeError):
|
||||
traceback.print_exc()
|
||||
|
||||
raise RuntimeError("Agent has failed to provide valid schema too many times")
|
||||
|
||||
def refine(self, structure: Structure) -> Structure:
|
||||
elements: list[StructureElement] = []
|
||||
context = _RefinerContext()
|
||||
|
||||
for window in self._windowizer.windowize(structure.elements):
|
||||
window.context = context.model_dump(mode="json")
|
||||
intermediate = self._refine_window(window)
|
||||
elements.extend(intermediate.elements)
|
||||
context = intermediate.context
|
||||
|
||||
return Structure(elements=elements)
|
||||
Reference in New Issue
Block a user