Huge refactoring
This commit is contained in:
111
asr_eventizer.py
Normal file
111
asr_eventizer.py
Normal file
@@ -0,0 +1,111 @@
|
||||
from dataclasses import asdict
|
||||
from typing import Any
|
||||
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"""
|
||||
|
||||
def _eventize_window(self, window: Window[AsrFilterSegment]) -> _EventizeResult:
|
||||
"""Eventize a single window"""
|
||||
messages = [
|
||||
self._system_prompt,
|
||||
AgentMessage(
|
||||
content=json.dumps(asdict(window), indent=2, ensure_ascii=False),
|
||||
role="user"
|
||||
)
|
||||
]
|
||||
retries_left = 5
|
||||
while retries_left > 0:
|
||||
retries_left -= 1
|
||||
response = self._agent.completion(messages=messages)
|
||||
# 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:
|
||||
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.json", "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
|
||||
# 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
|
||||
))
|
||||
# final timeline
|
||||
return Timeline(events=events)
|
||||
Reference in New Issue
Block a user