WIP
This commit is contained in:
@@ -1,5 +1,7 @@
|
||||
from dataclasses import asdict
|
||||
from typing import Any
|
||||
import traceback
|
||||
import shutil
|
||||
import os
|
||||
import json
|
||||
|
||||
from pydantic import BaseModel
|
||||
@@ -28,20 +30,32 @@ class _EventizeResult(BaseModel):
|
||||
|
||||
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(asdict(window), indent=2, ensure_ascii=False),
|
||||
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)
|
||||
@@ -56,6 +70,7 @@ class AsrEventizer:
|
||||
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"
|
||||
@@ -70,7 +85,7 @@ class AsrEventizer:
|
||||
"""
|
||||
self._agent = agent
|
||||
self._windowizer = windowizer
|
||||
with open("prompts/asr_eventizer.json", "r") as f:
|
||||
with open("prompts/asr_eventizer.md", "r") as f:
|
||||
self._system_prompt = AgentMessage(
|
||||
content=f.read(),
|
||||
role="system"
|
||||
@@ -91,12 +106,13 @@ class AsrEventizer:
|
||||
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]
|
||||
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(
|
||||
|
||||
Reference in New Issue
Block a user