Added reference resolver WIP (codex)

This commit is contained in:
2026-09-23 01:42:50 +03:00
parent 0da97ca225
commit 1e462b22a9
6 changed files with 564 additions and 11 deletions

365
reference_resolver.py Normal file
View File

@@ -0,0 +1,365 @@
import json
import logging
import math
import os
import shutil
import subprocess
import tempfile
import traceback
from dataclasses import dataclass
from pathlib import Path
from typing import Literal
from pydantic import BaseModel, ConfigDict, ValidationError, model_validator
from agent import Agent, AgentMessage
from utils import Event, Timeline
from video_references import UnresolvedReference, UnresolvedReferences
class _StrictModel(BaseModel):
model_config = ConfigDict(extra="forbid", strict=True)
class FailedReference(UnresolvedReference):
reason: Literal["no_frames", "not_found"]
class ReferenceResolutionResult(_StrictModel):
events: list[Event]
unresolved: list[FailedReference]
class _FrameResult(_StrictModel):
matched: bool
image_number: int | None
description: str | None
ocr_text: str | None
@model_validator(mode="after")
def validate_match(self) -> "_FrameResult":
if self.matched and (self.image_number is None or self.image_number < 1):
raise ValueError("A matched result must contain image_number")
if not self.matched and any(
value is not None
for value in (self.image_number, self.description, self.ocr_text)
):
raise ValueError("An unmatched result must not contain image data")
return self
@dataclass(frozen=True)
class _Frame:
timestamp: float
path: str
class ReferenceResolver:
"""Resolve visual references against sampled video frames."""
DEBUG_ID = 0
MAX_RETRIES = 3
def __init__(
self,
agent: Agent,
*,
video_path: str = "video.mp4",
image_dir: str = "images",
frame_interval: float = 1.0,
max_frames: int = 9,
) -> None:
if frame_interval <= 0:
raise ValueError("frame_interval must be positive")
if max_frames < 1:
raise ValueError("max_frames must be positive")
self._agent = agent
self._video_path = video_path
self._image_dir = Path(image_dir)
self._frame_interval = frame_interval
self._max_frames = max_frames
with open("prompts/reference_resolver.md", "r", encoding="utf-8") as f:
self._system_prompt = AgentMessage(content=f.read(), role="system")
def _sampling_timestamps(self, start: float, end: float) -> list[float]:
start = max(0.0, start)
end = max(start, end)
duration = end - start
if duration <= 0.05:
return [start]
count = min(
max(1, math.ceil(duration / self._frame_interval)),
self._max_frames,
)
step = duration / count
return [start + step * (index + 0.5) for index in range(count)]
def _extract_frames(
self,
timestamps: list[float],
output_dir: str,
) -> list[_Frame]:
frames: list[_Frame] = []
for index, timestamp in enumerate(timestamps):
output_path = os.path.join(output_dir, f"frame_{index:04d}.jpg")
command = [
"ffmpeg",
"-hide_banner",
"-loglevel", "error",
"-y",
"-ss", f"{timestamp:.3f}",
"-i", self._video_path,
"-frames:v", "1",
"-vf", "scale=1280:1280:force_original_aspect_ratio=decrease",
"-q:v", "3",
output_path,
]
try:
subprocess.run(command, check=True)
except subprocess.CalledProcessError:
logging.warning("Failed to extract frame at %.3f", timestamp)
continue
if os.path.isfile(output_path):
frames.append(_Frame(timestamp=timestamp, path=output_path))
return frames
@staticmethod
def _search_levels(frame_count: int) -> list[list[int]]:
if frame_count == 0:
return []
first = list(dict.fromkeys([0, (frame_count - 1) // 2, frame_count - 1]))
levels = [first]
selected = set(first)
while len(selected) < frame_count:
next_level: list[int] = []
ordered = sorted(selected)
for left, right in zip(ordered, ordered[1:]):
if right - left <= 1:
continue
midpoint = (left + right) // 2
if midpoint not in selected:
next_level.append(midpoint)
if not next_level:
break
levels.append(next_level)
selected.update(next_level)
return levels
@staticmethod
def _validate_frame_result(
result: _FrameResult,
reference: UnresolvedReference,
frame_count: int,
) -> None:
if not result.matched:
return
if result.image_number is None or result.image_number > frame_count:
raise ValueError("image_number does not exist in the current batch")
if reference.type == "vis":
if result.description is None or not result.description.strip():
raise ValueError("A visual match must contain description")
if result.ocr_text is not None:
raise ValueError("A visual match must not contain ocr_text")
else:
if result.ocr_text is None or not result.ocr_text.strip():
raise ValueError("An OCR match must contain ocr_text")
if result.description is not None:
raise ValueError("An OCR match must not contain description")
def _check_frames(
self,
reference: UnresolvedReference,
frames: list[_Frame],
) -> _FrameResult:
request = {
"reference": reference.model_dump(mode="json"),
"images": [
{
"image_number": index,
"timestamp": frame.timestamp,
}
for index, frame in enumerate(frames, start=1)
],
}
messages = [
self._system_prompt,
AgentMessage(
content=json.dumps(request, indent=2, ensure_ascii=False),
role="user",
),
]
debug_dir = None
if os.path.isdir("debug"):
debug_dir = f"debug/ReferenceResolver/{ReferenceResolver.DEBUG_ID}"
ReferenceResolver.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_with_images(
messages,
[frame.path for frame in frames],
detail="high" if reference.type == "ocr" else "low",
response_format=_FrameResult,
)
if debug_dir:
with open(
f"{debug_dir}/{retries_left}-retries-left.txt",
"w",
encoding="utf-8",
) as f:
f.write(response)
try:
result = _FrameResult.model_validate_json(response)
self._validate_frame_result(result, reference, len(frames))
return result
except (ValidationError, ValueError, TypeError):
traceback.print_exc()
raise RuntimeError("Agent has failed to provide valid schema too many times")
def _find_frame(
self,
reference: UnresolvedReference,
frames: list[_Frame],
) -> tuple[_Frame, _FrameResult] | None:
for level in self._search_levels(len(frames)):
batch = [frames[index] for index in level]
result = self._check_frames(reference, batch)
if result.matched:
assert result.image_number is not None
return batch[result.image_number - 1], result
return None
def _resolved_event(
self,
reference_index: int,
reference: UnresolvedReference,
frame: _Frame,
result: _FrameResult,
) -> Event:
if reference.type == "ocr":
assert result.ocr_text is not None
return Event(
id=-1,
type="ocr",
timestamp=frame.timestamp,
duration=0.0,
text=reference.text,
payload=result.ocr_text,
)
assert result.description is not None
self._image_dir.mkdir(parents=True, exist_ok=True)
timestamp_ms = round(frame.timestamp * 1000)
destination = self._image_dir / (
f"reference_{reference_index:04d}_{timestamp_ms:010d}.jpg"
)
shutil.copyfile(frame.path, destination)
return Event(
id=-1,
type="vis",
timestamp=frame.timestamp,
duration=0.0,
text=result.description,
payload=destination.as_posix(),
)
def resolve(
self,
timeline: Timeline,
references: UnresolvedReferences,
) -> ReferenceResolutionResult:
events_by_id: dict[int, Event] = {}
positions: dict[int, int] = {}
for position, event in enumerate(timeline.events):
if event.type != "voice":
raise ValueError("Reference resolver expects voice events")
if event.id in events_by_id:
raise ValueError(f"Duplicate event ID: {event.id}")
events_by_id[event.id] = event
positions[event.id] = position
resolved: list[Event] = []
failed: list[FailedReference] = []
frame_cache: dict[tuple[int, ...], list[_Frame]] = {}
with tempfile.TemporaryDirectory(prefix="sumka-frames-") as temp_dir:
for reference_index, reference in enumerate(references.unresolved):
if any(event_id not in events_by_id for event_id in reference.ids):
raise ValueError("Reference contains an unknown event ID")
reference_positions = [positions[event_id] for event_id in reference.ids]
if reference_positions != sorted(reference_positions):
raise ValueError("Reference ids must be in chronological order")
cache_key = tuple(reference.ids)
frames = frame_cache.get(cache_key)
if frames is None:
related = [events_by_id[event_id] for event_id in reference.ids]
start = min(event.timestamp for event in related)
end = max(
event.timestamp + max(0.0, event.duration)
for event in related
)
timestamps = self._sampling_timestamps(start, end)
reference_dir = os.path.join(temp_dir, f"reference_{reference_index:04d}")
os.makedirs(reference_dir, exist_ok=True)
frames = self._extract_frames(timestamps, reference_dir)
frame_cache[cache_key] = frames
if not frames:
failed.append(FailedReference(
**reference.model_dump(mode="json"),
reason="no_frames",
))
continue
match = self._find_frame(reference, frames)
if match is None:
failed.append(FailedReference(
**reference.model_dump(mode="json"),
reason="not_found",
))
continue
frame, result = match
resolved.append(
self._resolved_event(
reference_index,
reference,
frame,
result,
)
)
all_events = timeline.events + resolved
all_events.sort(key=lambda event: (event.timestamp, event.type != "voice"))
voice_ids: dict[int, int] = {
event.id: event_id
for event_id, event in enumerate(all_events)
if event.type == "voice"
}
final_events = [
event.model_copy(update={"id": event_id})
for event_id, event in enumerate(all_events)
]
final_failed = [
reference.model_copy(update={
"ids": [voice_ids[event_id] for event_id in reference.ids]
})
for reference in failed
]
return ReferenceResolutionResult(
events=final_events,
unresolved=final_failed,
)