500 lines
17 KiB
Python
500 lines
17 KiB
Python
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 paths import debug_path, prompt_path
|
|
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
|
|
crop_box: list[int] | 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.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
|
|
path: str
|
|
|
|
|
|
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,
|
|
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
|
|
self._system_prompt = AgentMessage(
|
|
content=prompt_path("reference_resolver.md").read_text(encoding="utf-8"),
|
|
role="system",
|
|
)
|
|
self._ocr_prompt = AgentMessage(
|
|
content=prompt_path("reference_resolver_ocr.md").read_text(encoding="utf-8"),
|
|
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.crop_box is not None:
|
|
raise ValueError("A visual match must not contain crop_box")
|
|
else:
|
|
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,
|
|
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 = debug_path("ReferenceResolver", ReferenceResolver.DEBUG_ID)
|
|
if debug_dir:
|
|
ReferenceResolver.DEBUG_ID += 1
|
|
(debug_dir / "request.txt").write_text(messages[1].content, encoding="utf-8")
|
|
|
|
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="low",
|
|
response_format=_FrameResult,
|
|
)
|
|
if debug_dir:
|
|
(debug_dir / f"{retries_left}-retries-left.txt").write_text(
|
|
response,
|
|
encoding="utf-8",
|
|
)
|
|
|
|
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 _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 = debug_path("ReferenceResolverOCR", ReferenceResolver.OCR_DEBUG_ID)
|
|
if debug_dir:
|
|
ReferenceResolver.OCR_DEBUG_ID += 1
|
|
(debug_dir / "request.txt").write_text(messages[1].content, encoding="utf-8")
|
|
shutil.copyfile(crop_path, 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:
|
|
(debug_dir / f"{retries_left}-retries-left.txt").write_text(
|
|
response,
|
|
encoding="utf-8",
|
|
)
|
|
|
|
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, 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 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(
|
|
self,
|
|
reference_index: int,
|
|
reference: UnresolvedReference,
|
|
frame: _Frame,
|
|
result: _FrameResult,
|
|
ocr_text: str | None,
|
|
) -> Event:
|
|
if reference.type == "ocr":
|
|
assert ocr_text is not None
|
|
return Event(
|
|
id=-1,
|
|
type="ocr",
|
|
timestamp=frame.timestamp,
|
|
duration=0.0,
|
|
text=reference.text,
|
|
payload=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, ocr_text = match
|
|
resolved.append(
|
|
self._resolved_event(
|
|
reference_index,
|
|
reference,
|
|
frame,
|
|
result,
|
|
ocr_text,
|
|
)
|
|
)
|
|
|
|
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,
|
|
)
|