Improved video referencing WIP (codex)

This commit is contained in:
2026-09-23 02:08:50 +03:00
parent 1e462b22a9
commit e5e09c00f9
8 changed files with 222 additions and 37 deletions

View File

@@ -34,7 +34,7 @@ class _FrameResult(_StrictModel):
matched: bool
image_number: int | None
description: str | None
ocr_text: str | None
crop_box: list[int] | None
@model_validator(mode="after")
def validate_match(self) -> "_FrameResult":
@@ -42,12 +42,25 @@ class _FrameResult(_StrictModel):
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)
for value in (self.image_number, self.description, self.crop_box)
):
raise ValueError("An unmatched result must not contain image data")
return self
class _OCRResult(_StrictModel):
matched: bool
ocr_text: str | None
@model_validator(mode="after")
def validate_match(self) -> "_OCRResult":
if self.matched and (self.ocr_text is None or not self.ocr_text.strip()):
raise ValueError("A matched OCR result must contain ocr_text")
if not self.matched and self.ocr_text is not None:
raise ValueError("An unmatched OCR result must not contain ocr_text")
return self
@dataclass(frozen=True)
class _Frame:
timestamp: float
@@ -58,7 +71,9 @@ class ReferenceResolver:
"""Resolve visual references against sampled video frames."""
DEBUG_ID = 0
OCR_DEBUG_ID = 0
MAX_RETRIES = 3
CROP_PADDING = 30
def __init__(
self,
@@ -81,6 +96,8 @@ class ReferenceResolver:
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")
with open("prompts/reference_resolver_ocr.md", "r", encoding="utf-8") as f:
self._ocr_prompt = AgentMessage(content=f.read(), role="system")
def _sampling_timestamps(self, start: float, end: float) -> list[float]:
start = max(0.0, start)
@@ -163,13 +180,19 @@ class ReferenceResolver:
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")
if result.crop_box is not None:
raise ValueError("A visual match must not contain crop_box")
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")
if result.crop_box is None or len(result.crop_box) != 4:
raise ValueError("An OCR match must contain crop_box")
y_min, x_min, y_max, x_max = result.crop_box
if not (
0 <= y_min < y_max <= 1000
and 0 <= x_min < x_max <= 1000
):
raise ValueError("crop_box must be within 0..1000")
def _check_frames(
self,
@@ -208,7 +231,7 @@ class ReferenceResolver:
response = self._agent.completion_with_images(
messages,
[frame.path for frame in frames],
detail="high" if reference.type == "ocr" else "low",
detail="low",
response_format=_FrameResult,
)
if debug_dir:
@@ -228,17 +251,131 @@ class ReferenceResolver:
raise RuntimeError("Agent has failed to provide valid schema too many times")
def _crop_frame(
self,
frame: _Frame,
crop_box: list[int],
) -> str | None:
y_min, x_min, y_max, x_max = crop_box
padding = self.CROP_PADDING
y_min = max(0, y_min - padding)
x_min = max(0, x_min - padding)
y_max = min(1000, y_max + padding)
x_max = min(1000, x_max + padding)
width = (x_max - x_min) / 1000
height = (y_max - y_min) / 1000
x = x_min / 1000
y = y_min / 1000
crop_path = str(Path(frame.path).with_name(f"{Path(frame.path).stem}_crop.jpg"))
crop_filter = (
f"crop=trunc(iw*{width:.6f}/2)*2:"
f"trunc(ih*{height:.6f}/2)*2:"
f"iw*{x:.6f}:ih*{y:.6f},"
"scale=1600:1600:force_original_aspect_ratio=decrease:"
"force_divisible_by=2"
)
command = [
"ffmpeg",
"-hide_banner",
"-loglevel", "error",
"-y",
"-ss", f"{frame.timestamp:.3f}",
"-i", self._video_path,
"-frames:v", "1",
"-vf", crop_filter,
"-q:v", "2",
crop_path,
]
try:
subprocess.run(command, check=True)
except subprocess.CalledProcessError:
logging.warning("Failed to crop frame at %.3f", frame.timestamp)
return None
return crop_path if os.path.isfile(crop_path) else None
def _recognize_ocr(
self,
reference: UnresolvedReference,
frame: _Frame,
crop_path: str,
) -> str | None:
request = {
"reference": reference.model_dump(mode="json"),
"timestamp": frame.timestamp,
}
messages = [
self._ocr_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/ReferenceResolverOCR/{ReferenceResolver.OCR_DEBUG_ID}"
ReferenceResolver.OCR_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)
shutil.copyfile(crop_path, f"{debug_dir}/crop.jpg")
retries_left = self.MAX_RETRIES
while retries_left > 0:
retries_left -= 1
response = self._agent.completion_with_images(
messages,
[crop_path],
detail="high",
response_format=_OCRResult,
)
if debug_dir:
with open(
f"{debug_dir}/{retries_left}-retries-left.txt",
"w",
encoding="utf-8",
) as f:
f.write(response)
try:
result = _OCRResult.model_validate_json(response)
if not result.matched:
return None
assert result.ocr_text is not None
return result.ocr_text.strip()
except (ValidationError, ValueError, TypeError):
traceback.print_exc()
raise RuntimeError("Agent has failed to provide valid OCR schema too many times")
def _find_frame(
self,
reference: UnresolvedReference,
frames: list[_Frame],
) -> tuple[_Frame, _FrameResult] | None:
) -> tuple[_Frame, _FrameResult, str | None] | 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
if not result.matched:
continue
assert result.image_number is not None
frame = batch[result.image_number - 1]
if reference.type == "vis":
return frame, result, None
assert result.crop_box is not None
crop_path = self._crop_frame(frame, result.crop_box)
if crop_path is None:
continue
ocr_text = self._recognize_ocr(reference, frame, crop_path)
if ocr_text is not None:
return frame, result, ocr_text
logging.info(
"OCR crop at %.3f is unreadable, continuing search",
frame.timestamp,
)
return None
def _resolved_event(
@@ -247,16 +384,17 @@ class ReferenceResolver:
reference: UnresolvedReference,
frame: _Frame,
result: _FrameResult,
ocr_text: str | None,
) -> Event:
if reference.type == "ocr":
assert result.ocr_text is not None
assert ocr_text is not None
return Event(
id=-1,
type="ocr",
timestamp=frame.timestamp,
duration=0.0,
text=reference.text,
payload=result.ocr_text,
payload=ocr_text,
)
assert result.description is not None
@@ -332,13 +470,14 @@ class ReferenceResolver:
))
continue
frame, result = match
frame, result, ocr_text = match
resolved.append(
self._resolved_event(
reference_index,
reference,
frame,
result,
ocr_text,
)
)