125 lines
4.5 KiB
Python
125 lines
4.5 KiB
Python
from typing import Any
|
|
import traceback
|
|
import json
|
|
|
|
from pydantic import BaseModel
|
|
from agent import Agent, AgentMessage
|
|
from windowizer import Windowizer, Window
|
|
from asr_filter import AsrFilterResult, AsrFilterSegment
|
|
from paths import debug_path, prompt_path
|
|
from utils import Timeline, Event
|
|
|
|
class _PreEvent(BaseModel):
|
|
"""Objects of this schema are returned by AI"""
|
|
|
|
ids: list[int]
|
|
"""IDs of merged segment"""
|
|
|
|
text: str
|
|
"""Text of the event after segments merging"""
|
|
|
|
class _EventizeResult(BaseModel):
|
|
"""Result of a single window eventizing, as returned by AI"""
|
|
|
|
preevents: list[_PreEvent]
|
|
"""PreEvents, as returned by AI"""
|
|
|
|
context: dict[str, Any]
|
|
"""Context, as returned by AI"""
|
|
|
|
class AsrEventizer:
|
|
"""This class creates a list of events from AsrFilterResult"""
|
|
DEBUG_ID = 0
|
|
|
|
def _eventize_window(self, window: Window[AsrFilterSegment]) -> _EventizeResult:
|
|
"""Eventize a single window"""
|
|
messages = [
|
|
self._system_prompt,
|
|
AgentMessage(
|
|
content=json.dumps(window.model_dump(mode="json"), indent=2, ensure_ascii=False),
|
|
role="user"
|
|
)
|
|
]
|
|
debug_dir = debug_path("AsrEventizer", AsrEventizer.DEBUG_ID)
|
|
if debug_dir:
|
|
AsrEventizer.DEBUG_ID += 1
|
|
(debug_dir / "request.txt").write_text(messages[1].content, encoding="utf-8")
|
|
retries_left = 5
|
|
while retries_left > 0:
|
|
retries_left -= 1
|
|
response = self._agent.completion(messages=messages, response_format=_EventizeResult)
|
|
if debug_dir:
|
|
(debug_dir / f"{retries_left}-retries-left.txt").write_text(
|
|
response,
|
|
encoding="utf-8",
|
|
)
|
|
# validate data
|
|
try:
|
|
response = json.loads(response)
|
|
obj = _EventizeResult(**response)
|
|
past_ids = [e.id for e in window.past]
|
|
present_ids = [e.id for e in window.present]
|
|
usable_ids = past_ids + present_ids
|
|
# find invalid IDs
|
|
for p in obj.preevents:
|
|
for id in p.ids:
|
|
if id not in usable_ids:
|
|
raise RuntimeError("Model has tried to use ID that was not provided")
|
|
return obj
|
|
except:
|
|
traceback.print_exc()
|
|
continue
|
|
raise RuntimeError(
|
|
"Agent has failed to provide valid schema too many times"
|
|
)
|
|
|
|
def __init__(self, agent: Agent, windowizer: Windowizer[AsrFilterSegment]) -> None:
|
|
"""Create the eventizer.
|
|
|
|
Args:
|
|
- agent - agent that will be used
|
|
- windowizer - windowizer to use
|
|
"""
|
|
self._agent = agent
|
|
self._windowizer = windowizer
|
|
self._system_prompt = AgentMessage(
|
|
content=prompt_path("asr_eventizer.md").read_text(encoding="utf-8"),
|
|
role="system",
|
|
)
|
|
|
|
def eventize(self, asr_filter_result: AsrFilterResult) -> Timeline:
|
|
"""Builds event timeline from `AsrFilterResult`. Resulting timeline
|
|
consists only of `voice` events.
|
|
"""
|
|
events = []
|
|
windows = self._windowizer.windowize(asr_filter_result.segments)
|
|
context = {}
|
|
output_id = 0
|
|
last_processed_ids: list[int] = []
|
|
for window in windows:
|
|
# call the agent
|
|
window.context = context
|
|
window.past = [e for e in window.past if e.id not in last_processed_ids]
|
|
intermediate = self._eventize_window(window)
|
|
context = intermediate.context
|
|
last_processed_ids = []
|
|
# process preevents
|
|
for preevent in intermediate.preevents:
|
|
related_segments = [
|
|
ev for ev in window.past + window.present if ev.id in preevent.ids
|
|
]
|
|
last_processed_ids += [ev.id for ev in related_segments]
|
|
min_time = min(s.start for s in related_segments)
|
|
max_time = max(s.end for s in related_segments)
|
|
events.append(Event(
|
|
id=output_id,
|
|
type="voice",
|
|
timestamp=min_time,
|
|
duration=max_time - min_time,
|
|
text=preevent.text,
|
|
payload=None
|
|
))
|
|
output_id += 1
|
|
# final timeline
|
|
return Timeline(events=events)
|