128 lines
4.5 KiB
Python
128 lines
4.5 KiB
Python
from typing import Any
|
|
import traceback
|
|
import shutil
|
|
import os
|
|
import json
|
|
|
|
from pydantic import BaseModel
|
|
from agent import Agent, AgentMessage
|
|
from windowizer import Windowizer, Window
|
|
from asr_filter import AsrFilterResult, AsrFilterSegment
|
|
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 = None
|
|
if os.path.isdir("debug"):
|
|
debug_dir = f"debug/AsrEventizer/{AsrEventizer.DEBUG_ID}"
|
|
AsrEventizer.DEBUG_ID += 1
|
|
os.makedirs(debug_dir, exist_ok=True)
|
|
if debug_dir:
|
|
with open(f"{debug_dir}/request.txt", "w") as f:
|
|
f.write(messages[1].content)
|
|
retries_left = 5
|
|
while retries_left > 0:
|
|
retries_left -= 1
|
|
response = self._agent.completion(messages=messages)
|
|
if debug_dir:
|
|
with open(f"{debug_dir}/{retries_left}-retries-left.txt", "w") as f:
|
|
f.write(response)
|
|
# 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
|
|
with open("prompts/asr_eventizer.md", "r") as f:
|
|
self._system_prompt = AgentMessage(
|
|
content=f.read(),
|
|
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) |