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

1
.gitignore vendored
View File

@@ -2,6 +2,7 @@ __pycache__/
.venv/ .venv/
runtime/ runtime/
images/
debug/ debug/
output.md output.md
*.json *.json

View File

@@ -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"
}
] ]
} }
``` ```

View File

@@ -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,13 +35,13 @@ 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,
@@ -42,3 +49,56 @@ class Agent:
**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
View File

@@ -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

View 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
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,
)