Improved video referencing WIP (codex)
This commit is contained in:
@@ -324,6 +324,9 @@ Telegram ботом, которому отправили видео).
|
||||
- Кадры проверяются мультимодальной моделью от общего к частному: сначала
|
||||
начало, середина и конец интервала, затем середины оставшихся промежутков.
|
||||
После первого подходящего кадра поиск прекращается.
|
||||
- Для `ocr` кадр сначала выбирается в режиме `low`, затем `ffmpeg` локально
|
||||
вырезает нужную область и только этот фрагмент отправляется в `high` для
|
||||
точного распознавания. Если фрагмент не читается, поиск продолжается.
|
||||
- Для `vis` выбранный кадр сохраняется в директории `images/`. Для `ocr` в
|
||||
событии сохраняется распознанный текст.
|
||||
- Если подходящий кадр не найден, ссылка сохраняется в корневом списке
|
||||
|
||||
23
agent.py
23
agent.py
@@ -2,8 +2,7 @@ import base64
|
||||
import mimetypes
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Literal, cast
|
||||
from pydantic import BaseModel
|
||||
from typing import Any, Literal, cast
|
||||
|
||||
import httpx2
|
||||
from openai import OpenAI
|
||||
@@ -31,7 +30,17 @@ class Agent:
|
||||
**kwargs
|
||||
)
|
||||
self._model = model
|
||||
|
||||
|
||||
@staticmethod
|
||||
def _raw_content(response: Any) -> str:
|
||||
"""Read message content without triggering SDK-side schema validation."""
|
||||
try:
|
||||
content = response.http_response.json()["choices"][0]["message"]["content"]
|
||||
except (AttributeError, KeyError, IndexError, TypeError) as error:
|
||||
raise RuntimeError("Agent returned an invalid response") from error
|
||||
if not isinstance(content, str):
|
||||
raise RuntimeError("Agent returned an empty response")
|
||||
return content
|
||||
|
||||
def completion(self, messages: list[AgentMessage], **kwargs) -> str:
|
||||
"""Generate a completion for specified messages."""
|
||||
@@ -43,12 +52,12 @@ class Agent:
|
||||
"content": m.content
|
||||
})
|
||||
)
|
||||
response = self._client.chat.completions.parse(
|
||||
response = self._client.chat.completions.with_raw_response.parse(
|
||||
model=self._model,
|
||||
messages=messages_raw,
|
||||
**kwargs
|
||||
)
|
||||
return response.choices[0].message.content # type: ignore
|
||||
return self._raw_content(response)
|
||||
|
||||
def completion_with_images(
|
||||
self,
|
||||
@@ -96,9 +105,9 @@ class Agent:
|
||||
"content": content,
|
||||
})
|
||||
|
||||
response = self._client.chat.completions.parse(
|
||||
response = self._client.chat.completions.with_raw_response.parse(
|
||||
model=self._model,
|
||||
messages=messages_raw,
|
||||
**kwargs,
|
||||
)
|
||||
return response.choices[0].message.content # type: ignore
|
||||
return self._raw_content(response)
|
||||
|
||||
@@ -242,7 +242,6 @@ CONTEXT
|
||||
`pending.ids` должны содержать только ID из `present`, которые НЕ были добавлены в `preevents`, потому что их продолжение находится за границей текущего окна.
|
||||
|
||||
`pending.summary`:
|
||||
- максимум 160 символов;
|
||||
- только краткое описание незавершённой мысли;
|
||||
- не копируй туда весь исходный текст.
|
||||
|
||||
@@ -336,4 +335,4 @@ CONTEXT
|
||||
8. Ты не пересказал и не суммаризировал речь.
|
||||
9. Ты не добавил сведений, которых нет в ASR-сегментах.
|
||||
10. `context` остаётся коротким.
|
||||
11. Ответ является валидным JSON.
|
||||
11. Ответ является валидным JSON.
|
||||
|
||||
@@ -15,19 +15,24 @@
|
||||
таблицу, рисунок, объект или демонстрацию;
|
||||
- в `description` кратко и фактически опиши полезное содержимое выбранного
|
||||
кадра;
|
||||
- верни `ocr_text: null`.
|
||||
- верни `crop_box: null`.
|
||||
|
||||
Для `ocr`:
|
||||
|
||||
- выбери один кадр, на котором нужный текст виден достаточно полно и чётко;
|
||||
- точно перепиши только относящийся к запросу текст или формулу;
|
||||
- не исправляй и не дополняй распознанное по собственным знаниям;
|
||||
- сохрани обозначения, индексы, знаки и порядок элементов;
|
||||
- пока не переписывай текст: выбери один кадр, на котором присутствуют все
|
||||
запрошенные элементы;
|
||||
- укажи охватывающую их прямоугольную область как `crop_box` в формате
|
||||
`[y_min, x_min, y_max, x_max]`, где координаты нормализованы от 0 до 1000;
|
||||
- область должна включать весь относящийся к запросу текст или формулу и
|
||||
небольшой окружающий контекст;
|
||||
- верни `description: null`.
|
||||
|
||||
Не считай совпадением кадр, если содержимое не видно, обрезано, слишком мелкое
|
||||
или лишь предположительно соответствует запросу. Не используй текст запроса,
|
||||
чтобы выдумать отсутствующее содержимое.
|
||||
Для `vis` не считай совпадением кадр, если содержимое не видно, обрезано,
|
||||
слишком мелкое или лишь предположительно соответствует запросу. Для `ocr`
|
||||
мелкий текст допустим: сейчас нужно уверенно найти содержащую его область, а
|
||||
читать её будет следующий этап после увеличения. Не выбирай `ocr`-кадр, если
|
||||
нужная область отсутствует, обрезана, закрыта или её нельзя уверенно найти.
|
||||
Не используй текст запроса, чтобы выдумать отсутствующее содержимое.
|
||||
|
||||
Если подходят несколько кадров, выбери самый полный и читаемый.
|
||||
|
||||
@@ -37,7 +42,7 @@
|
||||
"matched": true,
|
||||
"image_number": 2,
|
||||
"description": "Схема системы управления с входным и выходным векторами.",
|
||||
"ocr_text": null
|
||||
"crop_box": null
|
||||
}
|
||||
|
||||
Для успешного `ocr`:
|
||||
@@ -46,7 +51,7 @@
|
||||
"matched": true,
|
||||
"image_number": 1,
|
||||
"description": null,
|
||||
"ocr_text": "ρ(F_E, F̄) < ε"
|
||||
"crop_box": [180, 240, 720, 810]
|
||||
}
|
||||
|
||||
Если ни один кадр не подходит:
|
||||
@@ -55,7 +60,7 @@
|
||||
"matched": false,
|
||||
"image_number": null,
|
||||
"description": null,
|
||||
"ocr_text": null
|
||||
"crop_box": null
|
||||
}
|
||||
|
||||
Не добавляй Markdown, комментарии или текст до и после JSON.
|
||||
|
||||
31
prompts/reference_resolver_ocr.md
Normal file
31
prompts/reference_resolver_ocr.md
Normal file
@@ -0,0 +1,31 @@
|
||||
Ты выполняешь точное OCR одного фрагмента кадра видеолекции.
|
||||
|
||||
В пользовательском сообщении передан `reference` — полный запрос на текст или
|
||||
формулу. Единственное приложенное изображение — увеличенная область, выбранная
|
||||
на предыдущем этапе.
|
||||
|
||||
- Верни `matched: true` только если видны и читаются все элементы, требуемые в
|
||||
`reference.text`.
|
||||
- Перепиши только запрошенное содержимое без исправлений, дополнений и догадок.
|
||||
- Сохрани черты над символами, индексы, регистр, знаки неравенств, скобки и
|
||||
порядок элементов.
|
||||
- Математические выражения записывай в LaTeX, обычный текст оставляй обычным.
|
||||
- Не подменяй нечёткие обозначения более привычными по смыслу.
|
||||
- Если хотя бы существенная часть обрезана или неоднозначна, верни
|
||||
`matched: false`.
|
||||
|
||||
Успешный результат:
|
||||
|
||||
{
|
||||
"matched": true,
|
||||
"ocr_text": "\\rho(\\bar{F}_{э}, \\bar{F}) < \\varepsilon"
|
||||
}
|
||||
|
||||
Если точное распознавание невозможно:
|
||||
|
||||
{
|
||||
"matched": false,
|
||||
"ocr_text": null
|
||||
}
|
||||
|
||||
Верни ровно один JSON без Markdown, комментариев или текста вокруг него.
|
||||
@@ -478,7 +478,6 @@ CONTEXT
|
||||
- эти события не должны быть полностью оформлены в готовые элементы.
|
||||
|
||||
`pending.summary`:
|
||||
- максимум 180 символов;
|
||||
- краткая подсказка для следующей итерации;
|
||||
- не является частью итогового документа.
|
||||
|
||||
@@ -585,4 +584,4 @@ OUTPUT
|
||||
10. Заголовки не создаются заново только из-за начала нового окна.
|
||||
11. Ты не добавил фактов, отсутствующих во входных событиях.
|
||||
12. `context` остаётся коротким.
|
||||
13. Ответ является валидным JSON.
|
||||
13. Ответ является валидным JSON.
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@@ -100,7 +100,7 @@ class _Pending(_StrictModel):
|
||||
"""An unfinished thought carried over to the next window."""
|
||||
|
||||
event_ids: list[int] = Field(min_length=1)
|
||||
summary: str = Field(min_length=1, max_length=180)
|
||||
summary: str = Field(min_length=1)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_event_ids(self) -> "_Pending":
|
||||
|
||||
Reference in New Issue
Block a user