Huge refactoring
This commit is contained in:
1
.gitignore
vendored
1
.gitignore
vendored
@@ -2,6 +2,7 @@ __pycache__/
|
||||
.venv/
|
||||
runtime/
|
||||
|
||||
output.md
|
||||
*.json
|
||||
|
||||
*.mkv
|
||||
|
||||
12
agent.py
12
agent.py
@@ -1,7 +1,17 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Literal
|
||||
|
||||
import httpx2
|
||||
from openai import OpenAI
|
||||
|
||||
from utils import AgentMessage
|
||||
|
||||
@dataclass
|
||||
class AgentMessage:
|
||||
content: str
|
||||
"""Content of the message"""
|
||||
|
||||
role: Literal["system", "assistant", "user"]
|
||||
"""Who sent the message"""
|
||||
|
||||
class Agent:
|
||||
"""Perform operations with timeline events using OpenAI-compatible API"""
|
||||
|
||||
@@ -1,10 +1,34 @@
|
||||
from typing import Any
|
||||
from dataclasses import dataclass
|
||||
from pydantic import BaseModel
|
||||
|
||||
import whisper
|
||||
|
||||
from utils import TimelineEvent
|
||||
class AsrRawSegment(BaseModel):
|
||||
"""Segment produced by audio recognition engine"""
|
||||
|
||||
class Transcriber:
|
||||
start: float
|
||||
"""Start of the segment"""
|
||||
|
||||
end: float
|
||||
"""End of the segment"""
|
||||
|
||||
text: str
|
||||
"""Text of the segment"""
|
||||
|
||||
engine: dict[str, Any]
|
||||
"""Engine-related data"""
|
||||
|
||||
class AsrRawResult(BaseModel):
|
||||
"""Result of transcribing"""
|
||||
|
||||
engine: str
|
||||
"""Name of the engine that was used for transcribing"""
|
||||
|
||||
segments: list[AsrRawSegment]
|
||||
"""Segments produced by the engine"""
|
||||
|
||||
class Asr:
|
||||
"""This class performs transcription of the audio file."""
|
||||
@staticmethod
|
||||
def get_models_list() -> list[str]:
|
||||
@@ -27,7 +51,7 @@ class Transcriber:
|
||||
)
|
||||
self._model = whisper.load_model(model, **kwargs)
|
||||
|
||||
def transcribe(self, path: str, **kwargs) -> list[TimelineEvent]:
|
||||
def recognize(self, path: str, **kwargs) -> AsrRawResult:
|
||||
"""Transcribe audiofile. The operation will take a lot of time for large
|
||||
files.
|
||||
|
||||
@@ -36,25 +60,24 @@ class Transcriber:
|
||||
- **kwargs - passed to `transcribe()`
|
||||
|
||||
Returns:
|
||||
- list of timeline events you should use
|
||||
- result of transcribing
|
||||
"""
|
||||
raw_segments: list[dict]
|
||||
raw_segments = self._model.transcribe(path, **kwargs)["segments"] # type: ignore
|
||||
result: list[TimelineEvent] = []
|
||||
seg_id: int = 0
|
||||
result = AsrRawResult(
|
||||
engine="whisper",
|
||||
segments=[]
|
||||
)
|
||||
for raw_segment in raw_segments:
|
||||
ev = TimelineEvent(
|
||||
id = f"asr_{seg_id}",
|
||||
timestamp=float(raw_segment["start"]),
|
||||
duration=float(raw_segment["end"]) - float(raw_segment["start"]),
|
||||
payload=raw_segment["text"],
|
||||
custom={
|
||||
e = AsrRawSegment(
|
||||
start=float(raw_segment["start"]),
|
||||
end=float(raw_segment["end"]),
|
||||
text=str(raw_segment["text"]),
|
||||
engine={
|
||||
"whisper_temperature": float(raw_segment["temperature"]),
|
||||
"whisper_avg_logprob": float(raw_segment["avg_logprob"]),
|
||||
"whisper_no_speech_prob": float(raw_segment["no_speech_prob"])
|
||||
},
|
||||
links=[]
|
||||
}
|
||||
)
|
||||
result.append(ev)
|
||||
seg_id += 1
|
||||
result.segments.append(e)
|
||||
return result
|
||||
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)
|
||||
47
asr_filter.py
Normal file
47
asr_filter.py
Normal file
@@ -0,0 +1,47 @@
|
||||
from dataclasses import dataclass
|
||||
from asr import AsrRawResult
|
||||
from pydantic import BaseModel
|
||||
|
||||
class AsrFilterSegment(BaseModel):
|
||||
"""Recognized audio segment after cleanup."""
|
||||
|
||||
id: int
|
||||
"""Segment ID, unique within AsrResult"""
|
||||
|
||||
start: float
|
||||
"""Segment start time"""
|
||||
|
||||
end: float
|
||||
"""Segment end time"""
|
||||
|
||||
text: str
|
||||
"""Segment text after cleanup"""
|
||||
|
||||
class AsrFilterResult(BaseModel):
|
||||
"""Result of AsrFilter"""
|
||||
|
||||
segments: list[AsrFilterSegment]
|
||||
"""List of produced segments"""
|
||||
|
||||
class AsrFilter:
|
||||
"""This class performs filtering of raw ASR segments and produces events."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
def filter(self, asr_raw_result: AsrRawResult) -> AsrFilterResult:
|
||||
"""Filters raw ASR segments."""
|
||||
result = AsrFilterResult(
|
||||
segments=[]
|
||||
)
|
||||
i = 0
|
||||
for orig in asr_raw_result.segments:
|
||||
s = AsrFilterSegment(
|
||||
id=i,
|
||||
start=orig.start,
|
||||
end=orig.end,
|
||||
text=orig.text.strip()
|
||||
)
|
||||
i += 1
|
||||
result.segments.append(s)
|
||||
return result
|
||||
406
main.py
406
main.py
@@ -7,13 +7,31 @@ import json
|
||||
import time
|
||||
import os
|
||||
from dataclasses import asdict
|
||||
from enum import Enum
|
||||
from typing import Callable
|
||||
|
||||
import torch
|
||||
from asr import Asr, AsrRawResult
|
||||
from asr_filter import AsrFilter
|
||||
|
||||
from transcriber import Transcriber
|
||||
from agent import Agent
|
||||
from utils import TimelineEvent, AnalysisWindow, AgentMessage, TimelineProcessingResult
|
||||
from utils import ffmpeg_split_video, ffmpeg_to_mp3
|
||||
|
||||
ARGS: argparse.Namespace
|
||||
|
||||
class Step(Enum):
|
||||
MEDIA_SEPARATION = "media_separation"
|
||||
VOICE_RECOGNITION = "voice_recognition"
|
||||
ASR_FILTER = "asr_filter"
|
||||
ASR_EVENTS = "asr_events"
|
||||
VIDEO_REFERENCES = "video_references"
|
||||
REFERENCE_RESOLVER = "reference_resolver"
|
||||
STRUCTURE_BUILDER = "structure_builder"
|
||||
MARKDOWN_BUILDER = "markdown_builder"
|
||||
|
||||
#
|
||||
# Utility
|
||||
#
|
||||
def check_cuda() -> None:
|
||||
"""Checks if CUDA is available."""
|
||||
if not torch.cuda.is_available():
|
||||
@@ -28,33 +46,22 @@ def setup_arguments() -> argparse.Namespace:
|
||||
Returns:
|
||||
- argparse namespace
|
||||
"""
|
||||
voice_models = Transcriber.get_models_list()
|
||||
voice_models = Asr.get_models_list()
|
||||
|
||||
parser = argparse.ArgumentParser(
|
||||
prog="sumka",
|
||||
description="Summarizes large video/audio files into convenient format",
|
||||
)
|
||||
parser.add_argument("filename", type=str)
|
||||
parser.add_argument(
|
||||
"--voice-model",
|
||||
"--asr-model",
|
||||
choices=voice_models,
|
||||
default="turbo" if "turbo" in voice_models else voice_models[-1]
|
||||
)
|
||||
parser.add_argument(
|
||||
"--voice-language",
|
||||
"--asr-language",
|
||||
choices=["ru", "en"],
|
||||
default="ru"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--window-payload-size",
|
||||
type=int,
|
||||
default=1024
|
||||
)
|
||||
parser.add_argument(
|
||||
"--window-context-size",
|
||||
type=int,
|
||||
default=128
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ai-model",
|
||||
type=str,
|
||||
@@ -67,287 +74,124 @@ def setup_arguments() -> argparse.Namespace:
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ai-api-key",
|
||||
type=str,
|
||||
default="gdsfgds"
|
||||
type=str
|
||||
)
|
||||
parser.add_argument("-v", action='store_true')
|
||||
return parser.parse_args()
|
||||
|
||||
#
|
||||
# GENERIC
|
||||
# Workflow
|
||||
#
|
||||
def make_windows(events: list[TimelineEvent], context_symbols: int, payload_symbols: int) -> list[AnalysisWindow]:
|
||||
result: list[AnalysisWindow] = []
|
||||
window_start = 0
|
||||
while window_start < len(events):
|
||||
window = AnalysisWindow([], [], [])
|
||||
result.append(window)
|
||||
# build the window itself
|
||||
window_end = window_start + 1
|
||||
total_size = 0
|
||||
for event in events[window_start:]:
|
||||
window.modifiable.append(event)
|
||||
total_size += len(event.payload)
|
||||
if total_size >= payload_symbols:
|
||||
break
|
||||
window_end += 1
|
||||
# build readonly events before the window
|
||||
total_size = 0
|
||||
for event in reversed(events[:window_start]):
|
||||
window.before_readonly.insert(0, event)
|
||||
total_size += len(event.payload)
|
||||
if total_size >= context_symbols:
|
||||
break
|
||||
# build readonly event after the window
|
||||
total_size = 0
|
||||
for event in events[window_end:]:
|
||||
window.after_readonly.append(event)
|
||||
total_size += len(event.payload)
|
||||
if total_size >= context_symbols:
|
||||
break
|
||||
# prepare for the next window
|
||||
window_start = window_end
|
||||
return [r for r in result if len(r.modifiable)]
|
||||
def on_media_separation(current_step: Step, input_data: dict | None) -> tuple[Step | None, dict | None]:
|
||||
# there may be no data for this step
|
||||
if input_data:
|
||||
raise RuntimeError("There must be no input data for Media Separation")
|
||||
# do not split if there's `audio.mp3`
|
||||
if os.path.isfile(WORKFLOW_DATA[current_step][0]):
|
||||
logging.info("Skipping media separation")
|
||||
return (Step.VOICE_RECOGNITION, None)
|
||||
|
||||
# filesnames to look for
|
||||
VIDEO_INPUTS = [
|
||||
"input.mp4",
|
||||
"input.mkv",
|
||||
"input.avi"
|
||||
]
|
||||
AUDIO_INPUTS = [
|
||||
"input.mp3",
|
||||
"input.m4a",
|
||||
"input.wav"
|
||||
]
|
||||
# execute video split
|
||||
for v in VIDEO_INPUTS:
|
||||
if os.path.isfile(v):
|
||||
logging.info(f"Splitting {v} to audio.mp3 and video.mp4")
|
||||
ffmpeg_split_video(v)
|
||||
return (Step.VOICE_RECOGNITION, None)
|
||||
# execute audio conversion
|
||||
for a in AUDIO_INPUTS:
|
||||
if os.path.isfile(a):
|
||||
logging.info(f"Converting {a} to audio.mp3")
|
||||
ffmpeg_to_mp3(a)
|
||||
return (Step.VOICE_RECOGNITION, None)
|
||||
logging.error(f"No supported `input.*` files found")
|
||||
return (None, None)
|
||||
|
||||
def execute_timeline_request(res: TimelineProcessingResult, request: dict):
|
||||
req = request["req"]
|
||||
if req == "modify":
|
||||
id = request["id"]
|
||||
payload = request["payload"]
|
||||
valid = [e for e in res.events if e.id == id]
|
||||
if not len(valid):
|
||||
raise RuntimeError(f"AI tries to modify nonexistent event with ID `{id}`")
|
||||
valid[0].payload = payload
|
||||
logging.info(f"Updated `{id}`'s payload to `{payload}`")
|
||||
else:
|
||||
print(json.dumps(request, indent=2, ensure_ascii=False))
|
||||
def on_voice_recognition(current_step: Step, input_data: dict | None) -> tuple[Step | None, dict | None]:
|
||||
# do not perform recognition if output file exists
|
||||
if os.path.isfile(WORKFLOW_DATA[current_step][0]):
|
||||
logging.info("Skipping voice recognition")
|
||||
with open(WORKFLOW_DATA[current_step][0], "rb") as f:
|
||||
return (Step.ASR_FILTER, json.load(f))
|
||||
logging.info(f"Loading ASR model `{ARGS.asr_model}`, language `{ARGS.asr_language}`")
|
||||
asr = Asr(ARGS.asr_model)
|
||||
logging.info(f"Speech recognition...")
|
||||
result = asr.recognize("audio.mp3", language=ARGS.asr_language)
|
||||
logging.info(f"Speech recognition done")
|
||||
return (Step.ASR_FILTER, result.model_dump(mode="json"))
|
||||
|
||||
def build_document_structure(timeline: TimelineProcessingResult, args: argparse.Namespace):
|
||||
results_path = "document_structure.json"
|
||||
result = []
|
||||
# create the agent
|
||||
agent = Agent(
|
||||
model=args.ai_model,
|
||||
base_url=args.ai_base_url,
|
||||
api_key=args.ai_api_key
|
||||
)
|
||||
# create the system prompt
|
||||
with open("prompts/build_structure.md", "r") as f:
|
||||
system_prompt = AgentMessage(
|
||||
f.read(),
|
||||
"system"
|
||||
)
|
||||
windows = make_windows(timeline.events, args.window_context_size * 2, args.window_payload_size * 2)
|
||||
# context
|
||||
context = {}
|
||||
# process each window
|
||||
for window_id, window in enumerate(windows):
|
||||
# prepare request body
|
||||
req = {
|
||||
"before_readonly": [e.get_ai_dict() for e in window.before_readonly],
|
||||
"content": [e.get_ai_dict() for e in window.modifiable],
|
||||
"after_readonly": [e.get_ai_dict() for e in window.after_readonly],
|
||||
"document_context": context
|
||||
}
|
||||
msg = AgentMessage(
|
||||
content=json.dumps(req, indent=2, ensure_ascii=False),
|
||||
role="user"
|
||||
)
|
||||
# process the window
|
||||
success = False
|
||||
attempt = 1
|
||||
while not success:
|
||||
logging.info(f"Processing a window #{window_id + 1} (attempt #{attempt})...")
|
||||
# try to parse as JSON
|
||||
try:
|
||||
# call the AI
|
||||
response = json.loads(
|
||||
agent.completion(messages=[system_prompt, msg]))
|
||||
# process each request separately
|
||||
for block in response["blocks"]:
|
||||
result.append(block)
|
||||
context = response["new_document_context"]
|
||||
success = True
|
||||
except:
|
||||
attempt += 1
|
||||
logging.error("Failed, retrying")
|
||||
logging.debug(traceback.format_exc())
|
||||
with open(results_path, "w") as f:
|
||||
json.dump(
|
||||
result,
|
||||
f,
|
||||
indent=4,
|
||||
ensure_ascii=False
|
||||
)
|
||||
return result
|
||||
def on_asr_filter(current_step: Step, input_data: dict | None) -> tuple[Step | None, dict | None]:
|
||||
# do not filter if output file exists
|
||||
if os.path.isfile(WORKFLOW_DATA[current_step][0]):
|
||||
logging.info("Skipping ASR filter")
|
||||
with open(WORKFLOW_DATA[current_step][0], "rb") as f:
|
||||
return (Step.ASR_EVENTS, json.load(f))
|
||||
# bad request
|
||||
if input_data is None:
|
||||
logging.error("Can'f filter raw ASR ouput without input_data")
|
||||
return (None, None)
|
||||
logging.info("Filtering raw ASR output...")
|
||||
filter = AsrFilter()
|
||||
result = filter.filter(AsrRawResult(**input_data))
|
||||
return (Step.ASR_EVENTS, result.model_dump(mode="json"))
|
||||
|
||||
#
|
||||
# AUDIO
|
||||
# Main
|
||||
#
|
||||
def transcribe_audio(audio_path: str,
|
||||
args: argparse.Namespace) -> list[TimelineEvent]:
|
||||
"""Transcribes audio.
|
||||
|
||||
Args:
|
||||
- audio_path - path to the audio file
|
||||
- result_path - path to the resulting JSON file
|
||||
- args - arguments as returned by argsparse
|
||||
|
||||
Returns:
|
||||
- timeline events produced by ASR
|
||||
"""
|
||||
# check if file exists and just load it if it does
|
||||
result_path = "asr_events.json"
|
||||
if os.path.isfile(result_path):
|
||||
try:
|
||||
logging.info(
|
||||
f"Trying to load transcription data from {result_path}"
|
||||
)
|
||||
with open(result_path, "rb") as f:
|
||||
j = [TimelineEvent(**e) for e in json.load(f)]
|
||||
logging.info(f"Loaded transcription data from {result_path}")
|
||||
return j
|
||||
except:
|
||||
logging.debug(
|
||||
f"Could not load transcription data from {result_path}"
|
||||
)
|
||||
# actually transcribe
|
||||
logging.info(f"Transcribing {audio_path}...")
|
||||
logging.debug(f"Creating transcriber (using model `{args.voice_model}`)")
|
||||
t = Transcriber(args.voice_model)
|
||||
logging.debug(f"Creating the transcription...")
|
||||
events = t.transcribe(
|
||||
audio_path,
|
||||
language=args.voice_language
|
||||
)
|
||||
logging.debug(f"Saving to {result_path}")
|
||||
with open(result_path, "w") as f:
|
||||
f.write(json.dumps([asdict(e) for e in events], indent=4, ensure_ascii=False))
|
||||
logging.info(f"Done transcribing, timeline events produced: {len(events)}")
|
||||
return events
|
||||
WORKFLOW_DATA: dict[Step, tuple[str, Callable[[Step, dict | None], tuple[Step | None, dict | None]] | None]] = {
|
||||
Step.MEDIA_SEPARATION: ("audio.mp3", on_media_separation),
|
||||
Step.VOICE_RECOGNITION: ("asr_raw.json", on_voice_recognition),
|
||||
Step.ASR_FILTER: ("asr.json", on_asr_filter),
|
||||
Step.ASR_EVENTS: ("audio_events.json", None),
|
||||
Step.VIDEO_REFERENCES: ("unresolved.json", None),
|
||||
Step.REFERENCE_RESOLVER: ("events.json", None),
|
||||
Step.STRUCTURE_BUILDER: ("structure.json", None),
|
||||
Step.MARKDOWN_BUILDER: ("output.md", None)
|
||||
}
|
||||
"""Information about workflow.
|
||||
|
||||
def prepare_audio_windows(events: list[TimelineEvent],
|
||||
args: argparse.Namespace) -> list[AnalysisWindow]:
|
||||
"""Prepare list of windows which should be processed by LLM.
|
||||
|
||||
Args:
|
||||
- events - return value of `transcribe_audio`
|
||||
- args - arguments as returned by argsparse
|
||||
|
||||
Returns:
|
||||
- list of windows for LLM
|
||||
"""
|
||||
return make_windows(
|
||||
events,
|
||||
args.window_context_size,
|
||||
args.window_payload_size
|
||||
)
|
||||
|
||||
def process_audio_windows(windows: list[AnalysisWindow], args: argparse.Namespace) -> TimelineProcessingResult:
|
||||
# check if already processed
|
||||
results_path = "asr_proc_events.json"
|
||||
if os.path.isfile(results_path):
|
||||
try:
|
||||
logging.info(f"Loading processing result from {results_path}")
|
||||
with open(results_path, "rb") as f:
|
||||
j = json.load(f)
|
||||
result = TimelineProcessingResult(
|
||||
events=[TimelineEvent(**e) for e in j["events"]],
|
||||
desired_events=[TimelineEvent(**e) for e in j["desired_events"]]
|
||||
)
|
||||
return result
|
||||
except:
|
||||
logging.error(traceback.print_exc())
|
||||
# create the agent
|
||||
agent = Agent(
|
||||
model=args.ai_model,
|
||||
base_url=args.ai_base_url,
|
||||
api_key=args.ai_api_key
|
||||
)
|
||||
# create the system prompt
|
||||
with open("prompts/audio_window.md", "r") as f:
|
||||
system_prompt = AgentMessage(
|
||||
f.read(),
|
||||
"system"
|
||||
)
|
||||
# prepare the processing result
|
||||
all_events = [w.modifiable for w in windows]
|
||||
processing_result = TimelineProcessingResult(
|
||||
events=[item for sublist in all_events for item in sublist],
|
||||
desired_events=[]
|
||||
)
|
||||
# context for AI to remember previous iterations
|
||||
previous_self_context = {
|
||||
"_comment": "Use this object as you data storage for next iterations"
|
||||
}
|
||||
# process each window
|
||||
for window_id, window in enumerate(windows):
|
||||
# prepare request body
|
||||
req = {
|
||||
"before_readonly": [e.get_ai_dict() for e in window.before_readonly],
|
||||
"modifiable": [e.get_ai_dict() for e in window.modifiable],
|
||||
"after_readonly": [e.get_ai_dict() for e in window.after_readonly],
|
||||
"ai_custom_context": previous_self_context
|
||||
}
|
||||
msg = AgentMessage(
|
||||
content=json.dumps(req, indent=2, ensure_ascii=False),
|
||||
role="user"
|
||||
)
|
||||
# process the window
|
||||
success = False
|
||||
attempt = 1
|
||||
while not success:
|
||||
logging.info(f"Processing a window #{window_id + 1} (attempt #{attempt})...")
|
||||
# try to parse as JSON
|
||||
try:
|
||||
# call the AI
|
||||
response = json.loads(
|
||||
agent.completion(messages=[system_prompt, msg]))
|
||||
# process each request separately
|
||||
requests = response["requests"]
|
||||
for r in requests:
|
||||
execute_timeline_request(processing_result, r)
|
||||
previous_self_context = response["new_context"]
|
||||
success = True
|
||||
except:
|
||||
attempt += 1
|
||||
logging.error("Failed, retrying")
|
||||
logging.debug(traceback.format_exc())
|
||||
with open(results_path, "w") as f:
|
||||
json.dump(
|
||||
{
|
||||
"events": [asdict(e) for e in processing_result.events],
|
||||
"desired_events": [asdict(e) for e in processing_result.desired_events],
|
||||
},
|
||||
f,
|
||||
indent=4,
|
||||
ensure_ascii=False
|
||||
)
|
||||
return processing_result
|
||||
Step.CODE: (
|
||||
"path/to/result.json",
|
||||
(cur_step: Step, input_data: dict | None)
|
||||
-> (next_step: Step | None, output_data: dict | None)
|
||||
)
|
||||
"""
|
||||
|
||||
def main() -> None:
|
||||
"""Application entry point"""
|
||||
global ARGS
|
||||
check_cuda()
|
||||
args = setup_arguments()
|
||||
# setup the logger
|
||||
logging.basicConfig(level=logging.DEBUG if args.v else logging.INFO)
|
||||
# transcribe
|
||||
audio_events = transcribe_audio(
|
||||
args.filename,
|
||||
args
|
||||
)
|
||||
# prepare audio windows
|
||||
audio_windows = prepare_audio_windows(
|
||||
audio_events,
|
||||
args
|
||||
)
|
||||
# process audio window
|
||||
processing_result = process_audio_windows(
|
||||
audio_windows,
|
||||
args
|
||||
)
|
||||
# build the document
|
||||
build_document_structure(processing_result, args)
|
||||
ARGS = setup_arguments()
|
||||
logging.basicConfig(level=logging.DEBUG if ARGS.v else logging.INFO)
|
||||
step_to_do = Step.MEDIA_SEPARATION
|
||||
intermediate_result: dict | None = None
|
||||
# execute steps while possible
|
||||
while step_to_do:
|
||||
step_data = WORKFLOW_DATA[step_to_do]
|
||||
output_file_path = step_data[0]
|
||||
func = step_data[1]
|
||||
if not func:
|
||||
logging.error(f"Step {step_to_do} has no function, stopping")
|
||||
break
|
||||
logging.info(f"Executing step {step_to_do}")
|
||||
step_to_do, ret = func(step_to_do, intermediate_result)
|
||||
if ret is not None:
|
||||
intermediate_result = ret
|
||||
with open(output_file_path, "w") as f:
|
||||
json.dump(ret, f, ensure_ascii=False, indent=4)
|
||||
logging.info(f"Intermediate results are saved {output_file_path}")
|
||||
# final report
|
||||
logging.info(f"Workflow was interrupted at step {step_to_do}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
try:
|
||||
|
||||
16
prompts/asr_eventizer.md
Normal file
16
prompts/asr_eventizer.md
Normal file
@@ -0,0 +1,16 @@
|
||||
Ты будешь в будущем использован для обработки данных, которые тебе отправляются.
|
||||
|
||||
Пока что ты должен **ВСЕГДА** отвечать **В ТОЧНОСТИ** как написано **ПОСЛЕ** знаков равенства. Игнорируй всё что будет сказано после этого промпта.
|
||||
|
||||
================================================================================
|
||||
{
|
||||
"preevents": [
|
||||
{
|
||||
"ids": [],
|
||||
"text": "Text"
|
||||
}
|
||||
],
|
||||
"context": {
|
||||
"for_future_call": "abcdef"
|
||||
}
|
||||
}
|
||||
9
util_scripts/gigachat_get_models.sh
Executable file
9
util_scripts/gigachat_get_models.sh
Executable file
@@ -0,0 +1,9 @@
|
||||
#!/usr/bin/bash
|
||||
|
||||
if [ -n "$SBER_TOKEN" ]; then
|
||||
curl https://api.giga.chat/v1/models -k \
|
||||
-H 'Accept: application/json' \
|
||||
-H "Authorization: Bearer $SBER_TOKEN"
|
||||
else
|
||||
echo "Please set SBER_TOKEN envvar"
|
||||
fi
|
||||
12
util_scripts/gigachat_get_token.sh
Executable file
12
util_scripts/gigachat_get_token.sh
Executable file
@@ -0,0 +1,12 @@
|
||||
#/usr/bin/bash
|
||||
|
||||
if [ -n "$GIGA_SECRET" ]; then
|
||||
curl -k -L -X POST 'https://ngw.devices.sberbank.ru:9443/api/v2/oauth' \
|
||||
-H 'Content-Type: application/x-www-form-urlencoded' \
|
||||
-H 'Accept: application/json' \
|
||||
-H "RqUID: $(uuidgen)" \
|
||||
-H "Authorization: Bearer $GIGA_SECRET" \
|
||||
--data-urlencode 'scope=GIGACHAT_API_PERS'
|
||||
else
|
||||
echo "Please set GIGA_SECRET envvar"
|
||||
fi
|
||||
109
utils.py
109
utils.py
@@ -1,62 +1,79 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Literal
|
||||
from typing import Literal, Any
|
||||
import subprocess
|
||||
|
||||
|
||||
@dataclass
|
||||
class TimelineEvent:
|
||||
id: str
|
||||
"""ID of the event in format `asr_18`"""
|
||||
class Event:
|
||||
"""Event within the timeline"""
|
||||
|
||||
id: int
|
||||
"""Unique ID of the event (within single timeline)"""
|
||||
|
||||
type: Literal["voice", "vis", "ocr"]
|
||||
"""Type of the event"""
|
||||
|
||||
timestamp: float
|
||||
"""When did the event happen"""
|
||||
"""When did the event start (or happen, if `duration` is 0)"""
|
||||
|
||||
duration: float
|
||||
"""How long did the event last (zero if it does not make sence)"""
|
||||
"""How long did the event last (0 if can't apply)"""
|
||||
|
||||
payload: str
|
||||
"""Payload of the event (text for `asr`, description for `vis`, OCR result for `ocr`)"""
|
||||
text: str
|
||||
"""Textual representation of the event (or data in `payload`, if it is not None)"""
|
||||
|
||||
custom: dict
|
||||
"""Custom data"""
|
||||
|
||||
links: list[str]
|
||||
"""ID of related timeline events, empty list for None"""
|
||||
|
||||
def get_ai_dict(self) -> dict:
|
||||
"""Returns dict that is sanitized for AI."""
|
||||
d = {
|
||||
"id": self.id,
|
||||
"timestamp": self.timestamp,
|
||||
"payload": self.payload
|
||||
}
|
||||
if self.duration:
|
||||
d["duration"] = self.duration
|
||||
if self.links:
|
||||
d["links"] = self.links
|
||||
return d
|
||||
payload: str | None
|
||||
"""Payload of the event. `None` for `voice`, image path for `vis`, OCR result for `ocr`"""
|
||||
|
||||
@dataclass
|
||||
class AnalysisWindow:
|
||||
before_readonly: list[TimelineEvent]
|
||||
"""Timeline events for context (before current window)"""
|
||||
class Timeline:
|
||||
"""Timeline of events"""
|
||||
|
||||
modifiable: list[TimelineEvent]
|
||||
"""Timeline events that are to be analyzed"""
|
||||
events: list[Event]
|
||||
"""Events of the timeline"""
|
||||
|
||||
after_readonly: list[TimelineEvent]
|
||||
"""Timeline events for context (after current window)"""
|
||||
def ffmpeg_split_video(input_path: str) -> tuple[str, str]:
|
||||
"""Extract audio and video into ./audio.mp3 and ./video.mp4"""
|
||||
audio_path = "audio.mp3"
|
||||
video_path = "video.mp4"
|
||||
subprocess.run(
|
||||
[
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-i", input_path,
|
||||
|
||||
@dataclass
|
||||
class AgentMessage:
|
||||
content: str
|
||||
"""Content of the message"""
|
||||
# Audio
|
||||
"-map", "0:a:0",
|
||||
"-vn",
|
||||
"-c:a", "libmp3lame",
|
||||
"-q:a", "2",
|
||||
audio_path,
|
||||
|
||||
role: Literal["system", "assistant", "user"]
|
||||
"""Who sent the message"""
|
||||
# Video
|
||||
"-map", "0:v:0",
|
||||
"-an",
|
||||
"-c:v", "libx264",
|
||||
"-crf", "18",
|
||||
"-preset", "fast",
|
||||
video_path,
|
||||
],
|
||||
check=True,
|
||||
)
|
||||
return audio_path, video_path
|
||||
|
||||
@dataclass
|
||||
class TimelineProcessingResult:
|
||||
events: list[TimelineEvent]
|
||||
"""List of events"""
|
||||
|
||||
desired_events: list[TimelineEvent]
|
||||
"""List of events desired for existance"""
|
||||
def ffmpeg_to_mp3(input_path: str) -> str:
|
||||
"""Convert `input_path` path to mp3 and save it at ./audio.mp3"""
|
||||
output_path = "audio.mp3"
|
||||
subprocess.run(
|
||||
[
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-i", input_path,
|
||||
"-vn",
|
||||
"-c:a", "libmp3lame",
|
||||
"-q:a", "2",
|
||||
output_path,
|
||||
],
|
||||
check=True,
|
||||
)
|
||||
return output_path
|
||||
100
windowizer.py
Normal file
100
windowizer.py
Normal file
@@ -0,0 +1,100 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable, Iterable, Any
|
||||
|
||||
@dataclass
|
||||
class Window[T]:
|
||||
"""Window for AI processing."""
|
||||
|
||||
past: list[T]
|
||||
"""Part of data from previous window (can be empty)"""
|
||||
|
||||
present: list[T]
|
||||
"""Data for current window (can NOT be empty)"""
|
||||
|
||||
future: list[T]
|
||||
"""Part of data from the next window (can be empty)"""
|
||||
|
||||
context: dict[str, Any]
|
||||
"""Context that is preserved between LLM iterations"""
|
||||
|
||||
|
||||
class Windowizer[T]:
|
||||
"""This class builds windows from input data"""
|
||||
|
||||
@staticmethod
|
||||
def _build_list_with_limits(items: Iterable[T], min_count: int, max_size: int, size_func: Callable[[T], int]) -> list[T]:
|
||||
result: list[T] = []
|
||||
total_size: int = 0
|
||||
for i in items:
|
||||
result.append(i)
|
||||
total_size += size_func(i)
|
||||
if len(result) < min_count:
|
||||
continue
|
||||
if total_size >= max_size:
|
||||
break
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def default_size_func(elem: T) -> int:
|
||||
"""Default size function just returns length of the `repr` result."""
|
||||
return len(repr(elem))
|
||||
|
||||
def __init__(self,
|
||||
*,
|
||||
main_min_count: int = 2,
|
||||
main_max_size: int = 1024,
|
||||
side_min_count: int = 2,
|
||||
side_max_size: int = 256,
|
||||
size_func: Callable[[T], int] = default_size_func) -> None:
|
||||
"""Create Windowizer.
|
||||
|
||||
Args:
|
||||
- main_min_count - minimum count of items in `present`
|
||||
- main_max_size - maximum total size of items in `present`
|
||||
- side_min_count - minimum count of items in `past` and `future`
|
||||
- side_max_size - maximum total size of items in `past` and `future`
|
||||
- size_func - function that will be used to get item size
|
||||
"""
|
||||
self._main_min_count = main_min_count
|
||||
self._main_max_size = main_max_size
|
||||
self._side_min_count = side_min_count
|
||||
self._side_max_size = side_max_size
|
||||
self._size_func = size_func
|
||||
|
||||
def windowize(self, items: list[T]) -> list[Window[T]]:
|
||||
"""Builds windows from items.
|
||||
|
||||
`min_count_*` has more priority than `max_size_*`. It would be possible
|
||||
to build windows of 0 items otherwise.
|
||||
"""
|
||||
result: list[Window[T]] = []
|
||||
# create windows
|
||||
window_start = 0
|
||||
while window_start < len(items):
|
||||
window = Window[T](past=[], present=[], future=[], context={})
|
||||
# build `present`
|
||||
window.present = self._build_list_with_limits(
|
||||
items[window_start:],
|
||||
min_count=self._main_min_count,
|
||||
max_size=self._main_max_size,
|
||||
size_func=self._size_func
|
||||
)
|
||||
# build `past`
|
||||
window.past = self._build_list_with_limits(
|
||||
reversed(items[:window_start]),
|
||||
min_count=self._side_min_count,
|
||||
max_size=self._side_max_size,
|
||||
size_func=self._size_func
|
||||
)
|
||||
window.past.reverse()
|
||||
# build `future`
|
||||
window.future = self._build_list_with_limits(
|
||||
items[window_start+len(window.present):],
|
||||
min_count=self._side_min_count,
|
||||
max_size=self._side_max_size,
|
||||
size_func=self._size_func
|
||||
)
|
||||
window_start += len(window.present)
|
||||
result.append(window)
|
||||
# resulting windows
|
||||
return result
|
||||
Reference in New Issue
Block a user