From 1e462b22a90f8faa408d6a65ba4be034510e29c5 Mon Sep 17 00:00:00 2001 From: nikita Date: Wed, 23 Sep 2026 01:42:50 +0300 Subject: [PATCH] Added reference resolver WIP (codex) --- .gitignore | 1 + README.md | 25 ++- agent.py | 70 ++++++- main.py | 53 ++++- prompts/reference_resolver.md | 61 ++++++ reference_resolver.py | 365 ++++++++++++++++++++++++++++++++++ 6 files changed, 564 insertions(+), 11 deletions(-) create mode 100644 prompts/reference_resolver.md create mode 100644 reference_resolver.py diff --git a/.gitignore b/.gitignore index 87f6c6c..29e1a43 100644 --- a/.gitignore +++ b/.gitignore @@ -2,6 +2,7 @@ __pycache__/ .venv/ runtime/ +images/ debug/ output.md *.json diff --git a/README.md b/README.md index fd91a18..8c433e9 100644 --- a/README.md +++ b/README.md @@ -318,7 +318,22 @@ Telegram ботом, которому отправили видео). 6. **Разрешить ссылки на видео** - После этого шага должен быть получен файл `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 никак не связаны с предудущими шагами. ```json @@ -356,6 +371,14 @@ Telegram ботом, которому отправили видео). "text": "Для типа `voice` ключ `payload` всегда `null`", "payload": null } + ], + "unresolved": [ + { + "ids": [0, 3], + "type": "ocr", + "text": "Распознать формулу со слайда", + "reason": "not_found" + } ] } ``` diff --git a/agent.py b/agent.py index 379831c..43edc5e 100644 --- a/agent.py +++ b/agent.py @@ -1,9 +1,16 @@ +import base64 +import mimetypes from dataclasses import dataclass -from typing import Literal +from pathlib import Path +from typing import Literal, cast from pydantic import BaseModel import httpx2 from openai import OpenAI +from openai.types.chat import ( + ChatCompletionContentPartParam, + ChatCompletionMessageParam, +) @dataclass @@ -28,17 +35,70 @@ class Agent: def completion(self, messages: list[AgentMessage], **kwargs) -> str: """Generate a completion for specified messages.""" - messages_raw = [] + messages_raw: list[ChatCompletionMessageParam] = [] for m in messages: messages_raw.append( - { + cast(ChatCompletionMessageParam, { "role": m.role, "content": m.content - } + }) ) response = self._client.chat.completions.parse( model=self._model, messages=messages_raw, **kwargs ) - return response.choices[0].message.content # type: ignore \ No newline at end of file + 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 diff --git a/main.py b/main.py index b00a1e4..9d5fbd2 100644 --- a/main.py +++ b/main.py @@ -14,7 +14,8 @@ import torch from asr import Asr, AsrRawResult from asr_filter import AsrFilter, AsrFilterResult 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_refiner import StructureRefiner from markdown_builder import MarkdownBuilder @@ -100,6 +101,31 @@ def setup_arguments() -> argparse.Namespace: type=str, 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( "--structure-ai-model", 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)) if not os.path.isfile("video.mp4"): logging.info("No video, skipping") - return (Step.STRUCTURE_BUILDER, input_data) - # NOT IMPLEMENTED - logging.warning("Reference resolver is not implemented yet") - return (None, None) + with open("audio_events.json", "rb") as f: + return (Step.STRUCTURE_BUILDER, json.load(f)) + if input_data is 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]: # don't if done diff --git a/prompts/reference_resolver.md b/prompts/reference_resolver.md new file mode 100644 index 0000000..8d27d51 --- /dev/null +++ b/prompts/reference_resolver.md @@ -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. diff --git a/reference_resolver.py b/reference_resolver.py new file mode 100644 index 0000000..6006563 --- /dev/null +++ b/reference_resolver.py @@ -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, + )