Added reference resolver WIP (codex)
This commit is contained in:
1
.gitignore
vendored
1
.gitignore
vendored
@@ -2,6 +2,7 @@ __pycache__/
|
|||||||
.venv/
|
.venv/
|
||||||
runtime/
|
runtime/
|
||||||
|
|
||||||
|
images/
|
||||||
debug/
|
debug/
|
||||||
output.md
|
output.md
|
||||||
*.json
|
*.json
|
||||||
|
|||||||
25
README.md
25
README.md
@@ -318,7 +318,22 @@ Telegram ботом, которому отправили видео).
|
|||||||
6. **Разрешить ссылки на видео**
|
6. **Разрешить ссылки на видео**
|
||||||
- После этого шага должен быть получен файл `events.json`, используя который
|
- После этого шага должен быть получен файл `events.json`, используя который
|
||||||
будет построена структура итогового документа.
|
будет построена структура итогового документа.
|
||||||
- TODO, пока что это не реализовано
|
- Для каждого объекта из `unresolved.json` программа получает кадры через
|
||||||
|
`ffmpeg`. По умолчанию кадры берутся с интервалом в одну секунду, но не
|
||||||
|
более 9 кадров на одну ссылку.
|
||||||
|
- Кадры проверяются мультимодальной моделью от общего к частному: сначала
|
||||||
|
начало, середина и конец интервала, затем середины оставшихся промежутков.
|
||||||
|
После первого подходящего кадра поиск прекращается.
|
||||||
|
- Для `vis` выбранный кадр сохраняется в директории `images/`. Для `ocr` в
|
||||||
|
событии сохраняется распознанный текст.
|
||||||
|
- Если подходящий кадр не найден, ссылка сохраняется в корневом списке
|
||||||
|
`unresolved` файла `events.json` с причиной `not_found` или `no_frames`.
|
||||||
|
Её `ids` ссылаются на голосовые события из итогового списка `events`.
|
||||||
|
- Настройки модели задаются через `--resolver-ai-model`,
|
||||||
|
`--resolver-ai-base-url` и `--resolver-ai-api-key`. Период и максимальное
|
||||||
|
число кадров задаются через `--resolver-frame-interval` и
|
||||||
|
`--resolver-max-frames`. По умолчанию используется
|
||||||
|
`google/gemini-3.5-flash`.
|
||||||
- Итоговый файл `events.json`, должен иметь следующую структуру. ID никак
|
- Итоговый файл `events.json`, должен иметь следующую структуру. ID никак
|
||||||
не связаны с предудущими шагами.
|
не связаны с предудущими шагами.
|
||||||
```json
|
```json
|
||||||
@@ -356,6 +371,14 @@ Telegram ботом, которому отправили видео).
|
|||||||
"text": "Для типа `voice` ключ `payload` всегда `null`",
|
"text": "Для типа `voice` ключ `payload` всегда `null`",
|
||||||
"payload": null
|
"payload": null
|
||||||
}
|
}
|
||||||
|
],
|
||||||
|
"unresolved": [
|
||||||
|
{
|
||||||
|
"ids": [0, 3],
|
||||||
|
"type": "ocr",
|
||||||
|
"text": "Распознать формулу со слайда",
|
||||||
|
"reason": "not_found"
|
||||||
|
}
|
||||||
]
|
]
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|||||||
70
agent.py
70
agent.py
@@ -1,9 +1,16 @@
|
|||||||
|
import base64
|
||||||
|
import mimetypes
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Literal
|
from pathlib import Path
|
||||||
|
from typing import Literal, cast
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
|
|
||||||
import httpx2
|
import httpx2
|
||||||
from openai import OpenAI
|
from openai import OpenAI
|
||||||
|
from openai.types.chat import (
|
||||||
|
ChatCompletionContentPartParam,
|
||||||
|
ChatCompletionMessageParam,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -28,17 +35,70 @@ class Agent:
|
|||||||
|
|
||||||
def completion(self, messages: list[AgentMessage], **kwargs) -> str:
|
def completion(self, messages: list[AgentMessage], **kwargs) -> str:
|
||||||
"""Generate a completion for specified messages."""
|
"""Generate a completion for specified messages."""
|
||||||
messages_raw = []
|
messages_raw: list[ChatCompletionMessageParam] = []
|
||||||
for m in messages:
|
for m in messages:
|
||||||
messages_raw.append(
|
messages_raw.append(
|
||||||
{
|
cast(ChatCompletionMessageParam, {
|
||||||
"role": m.role,
|
"role": m.role,
|
||||||
"content": m.content
|
"content": m.content
|
||||||
}
|
})
|
||||||
)
|
)
|
||||||
response = self._client.chat.completions.parse(
|
response = self._client.chat.completions.parse(
|
||||||
model=self._model,
|
model=self._model,
|
||||||
messages=messages_raw,
|
messages=messages_raw,
|
||||||
**kwargs
|
**kwargs
|
||||||
)
|
)
|
||||||
return response.choices[0].message.content # type: ignore
|
return response.choices[0].message.content # type: ignore
|
||||||
|
|
||||||
|
def completion_with_images(
|
||||||
|
self,
|
||||||
|
messages: list[AgentMessage],
|
||||||
|
image_paths: list[str],
|
||||||
|
*,
|
||||||
|
detail: Literal["low", "high", "auto"] = "auto",
|
||||||
|
**kwargs,
|
||||||
|
) -> str:
|
||||||
|
"""Generate a completion with local images attached to the last message."""
|
||||||
|
if not messages or messages[-1].role != "user":
|
||||||
|
raise ValueError("The last message must be a user message")
|
||||||
|
if not image_paths:
|
||||||
|
raise ValueError("At least one image is required")
|
||||||
|
|
||||||
|
messages_raw: list[ChatCompletionMessageParam] = [
|
||||||
|
cast(ChatCompletionMessageParam, {
|
||||||
|
"role": message.role,
|
||||||
|
"content": message.content,
|
||||||
|
})
|
||||||
|
for message in messages[:-1]
|
||||||
|
]
|
||||||
|
content: list[ChatCompletionContentPartParam] = [
|
||||||
|
{"type": "text", "text": messages[-1].content}
|
||||||
|
]
|
||||||
|
for index, image_path in enumerate(image_paths, start=1):
|
||||||
|
path = Path(image_path)
|
||||||
|
mime_type = mimetypes.guess_type(path.name)[0] or "image/jpeg"
|
||||||
|
if not mime_type.startswith("image/"):
|
||||||
|
raise ValueError(f"Unsupported image type: {image_path}")
|
||||||
|
encoded = base64.b64encode(path.read_bytes()).decode("ascii")
|
||||||
|
content.append({
|
||||||
|
"type": "text",
|
||||||
|
"text": f"Изображение {index}",
|
||||||
|
})
|
||||||
|
content.append({
|
||||||
|
"type": "image_url",
|
||||||
|
"image_url": {
|
||||||
|
"url": f"data:{mime_type};base64,{encoded}",
|
||||||
|
"detail": detail,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
messages_raw.append({
|
||||||
|
"role": "user",
|
||||||
|
"content": content,
|
||||||
|
})
|
||||||
|
|
||||||
|
response = self._client.chat.completions.parse(
|
||||||
|
model=self._model,
|
||||||
|
messages=messages_raw,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
return response.choices[0].message.content # type: ignore
|
||||||
|
|||||||
53
main.py
53
main.py
@@ -14,7 +14,8 @@ import torch
|
|||||||
from asr import Asr, AsrRawResult
|
from asr import Asr, AsrRawResult
|
||||||
from asr_filter import AsrFilter, AsrFilterResult
|
from asr_filter import AsrFilter, AsrFilterResult
|
||||||
from asr_eventizer import AsrEventizer
|
from asr_eventizer import AsrEventizer
|
||||||
from video_references import VideoReferenceBuilder
|
from video_references import UnresolvedReferences, VideoReferenceBuilder
|
||||||
|
from reference_resolver import ReferenceResolver
|
||||||
from structure_builder import Structure, StructureBuilder
|
from structure_builder import Structure, StructureBuilder
|
||||||
from structure_refiner import StructureRefiner
|
from structure_refiner import StructureRefiner
|
||||||
from markdown_builder import MarkdownBuilder
|
from markdown_builder import MarkdownBuilder
|
||||||
@@ -100,6 +101,31 @@ def setup_arguments() -> argparse.Namespace:
|
|||||||
type=str,
|
type=str,
|
||||||
default=ai_api_key
|
default=ai_api_key
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--resolver-ai-model",
|
||||||
|
type=str,
|
||||||
|
default="google/gemini-3.5-flash"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--resolver-ai-base-url",
|
||||||
|
type=str,
|
||||||
|
default="https://api.proxyapi.ru/v1"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--resolver-ai-api-key",
|
||||||
|
type=str,
|
||||||
|
default=ai_api_key
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--resolver-frame-interval",
|
||||||
|
type=float,
|
||||||
|
default=1.0
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--resolver-max-frames",
|
||||||
|
type=int,
|
||||||
|
default=9
|
||||||
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--structure-ai-model",
|
"--structure-ai-model",
|
||||||
type=str,
|
type=str,
|
||||||
@@ -253,10 +279,27 @@ def on_reference_resolver(current_step: Step, input_data: dict | None) -> tuple[
|
|||||||
return (Step.STRUCTURE_BUILDER, json.load(f))
|
return (Step.STRUCTURE_BUILDER, json.load(f))
|
||||||
if not os.path.isfile("video.mp4"):
|
if not os.path.isfile("video.mp4"):
|
||||||
logging.info("No video, skipping")
|
logging.info("No video, skipping")
|
||||||
return (Step.STRUCTURE_BUILDER, input_data)
|
with open("audio_events.json", "rb") as f:
|
||||||
# NOT IMPLEMENTED
|
return (Step.STRUCTURE_BUILDER, json.load(f))
|
||||||
logging.warning("Reference resolver is not implemented yet")
|
if input_data is None:
|
||||||
return (None, None)
|
logging.error("Can't resolve video references without input_data")
|
||||||
|
return (None, None)
|
||||||
|
with open("audio_events.json", "rb") as f:
|
||||||
|
timeline = Timeline(**json.load(f))
|
||||||
|
logging.info("Creating the agent")
|
||||||
|
agent = Agent(
|
||||||
|
model=ARGS.resolver_ai_model,
|
||||||
|
base_url=ARGS.resolver_ai_base_url,
|
||||||
|
api_key=ARGS.resolver_ai_api_key
|
||||||
|
)
|
||||||
|
logging.info("Resolving video references...")
|
||||||
|
resolver = ReferenceResolver(
|
||||||
|
agent,
|
||||||
|
frame_interval=ARGS.resolver_frame_interval,
|
||||||
|
max_frames=ARGS.resolver_max_frames,
|
||||||
|
)
|
||||||
|
result = resolver.resolve(timeline, UnresolvedReferences(**input_data))
|
||||||
|
return (Step.STRUCTURE_BUILDER, result.model_dump(mode="json"))
|
||||||
|
|
||||||
def on_structure_builder(current_step: Step, input_data: dict | None) -> tuple[Step | None, dict | None]:
|
def on_structure_builder(current_step: Step, input_data: dict | None) -> tuple[Step | None, dict | None]:
|
||||||
# don't if done
|
# don't if done
|
||||||
|
|||||||
61
prompts/reference_resolver.md
Normal file
61
prompts/reference_resolver.md
Normal file
@@ -0,0 +1,61 @@
|
|||||||
|
Ты проверяешь кадры видеолекции и разрешаешь один запрос типа `vis` или `ocr`.
|
||||||
|
|
||||||
|
В пользовательском сообщении переданы:
|
||||||
|
|
||||||
|
- `reference` — что требуется найти;
|
||||||
|
- `images` — номера и временные метки приложенных изображений.
|
||||||
|
|
||||||
|
Изображения приложены после JSON в том же порядке и подписаны как
|
||||||
|
«Изображение 1», «Изображение 2» и так далее.
|
||||||
|
|
||||||
|
Для `vis`:
|
||||||
|
|
||||||
|
- выбери один кадр, который лучше всего соответствует запросу;
|
||||||
|
- кадр должен содержать действительно полезную схему, график, диаграмму,
|
||||||
|
таблицу, рисунок, объект или демонстрацию;
|
||||||
|
- в `description` кратко и фактически опиши полезное содержимое выбранного
|
||||||
|
кадра;
|
||||||
|
- верни `ocr_text: null`.
|
||||||
|
|
||||||
|
Для `ocr`:
|
||||||
|
|
||||||
|
- выбери один кадр, на котором нужный текст виден достаточно полно и чётко;
|
||||||
|
- точно перепиши только относящийся к запросу текст или формулу;
|
||||||
|
- не исправляй и не дополняй распознанное по собственным знаниям;
|
||||||
|
- сохрани обозначения, индексы, знаки и порядок элементов;
|
||||||
|
- верни `description: null`.
|
||||||
|
|
||||||
|
Не считай совпадением кадр, если содержимое не видно, обрезано, слишком мелкое
|
||||||
|
или лишь предположительно соответствует запросу. Не используй текст запроса,
|
||||||
|
чтобы выдумать отсутствующее содержимое.
|
||||||
|
|
||||||
|
Если подходят несколько кадров, выбери самый полный и читаемый.
|
||||||
|
|
||||||
|
Верни ровно один JSON:
|
||||||
|
|
||||||
|
{
|
||||||
|
"matched": true,
|
||||||
|
"image_number": 2,
|
||||||
|
"description": "Схема системы управления с входным и выходным векторами.",
|
||||||
|
"ocr_text": null
|
||||||
|
}
|
||||||
|
|
||||||
|
Для успешного `ocr`:
|
||||||
|
|
||||||
|
{
|
||||||
|
"matched": true,
|
||||||
|
"image_number": 1,
|
||||||
|
"description": null,
|
||||||
|
"ocr_text": "ρ(F_E, F̄) < ε"
|
||||||
|
}
|
||||||
|
|
||||||
|
Если ни один кадр не подходит:
|
||||||
|
|
||||||
|
{
|
||||||
|
"matched": false,
|
||||||
|
"image_number": null,
|
||||||
|
"description": null,
|
||||||
|
"ocr_text": null
|
||||||
|
}
|
||||||
|
|
||||||
|
Не добавляй Markdown, комментарии или текст до и после JSON.
|
||||||
365
reference_resolver.py
Normal file
365
reference_resolver.py
Normal 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,
|
||||||
|
)
|
||||||
Reference in New Issue
Block a user