Huge refactoring
This commit is contained in:
1
.gitignore
vendored
1
.gitignore
vendored
@@ -2,6 +2,7 @@ __pycache__/
|
|||||||
.venv/
|
.venv/
|
||||||
runtime/
|
runtime/
|
||||||
|
|
||||||
|
output.md
|
||||||
*.json
|
*.json
|
||||||
|
|
||||||
*.mkv
|
*.mkv
|
||||||
|
|||||||
12
agent.py
12
agent.py
@@ -1,7 +1,17 @@
|
|||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Literal
|
||||||
|
|
||||||
import httpx2
|
import httpx2
|
||||||
from openai import OpenAI
|
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:
|
class Agent:
|
||||||
"""Perform operations with timeline events using OpenAI-compatible API"""
|
"""Perform operations with timeline events using OpenAI-compatible API"""
|
||||||
|
|||||||
@@ -1,10 +1,34 @@
|
|||||||
from typing import Any
|
from typing import Any
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from pydantic import BaseModel
|
||||||
|
|
||||||
import whisper
|
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."""
|
"""This class performs transcription of the audio file."""
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def get_models_list() -> list[str]:
|
def get_models_list() -> list[str]:
|
||||||
@@ -27,7 +51,7 @@ class Transcriber:
|
|||||||
)
|
)
|
||||||
self._model = whisper.load_model(model, **kwargs)
|
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
|
"""Transcribe audiofile. The operation will take a lot of time for large
|
||||||
files.
|
files.
|
||||||
|
|
||||||
@@ -36,25 +60,24 @@ class Transcriber:
|
|||||||
- **kwargs - passed to `transcribe()`
|
- **kwargs - passed to `transcribe()`
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
- list of timeline events you should use
|
- result of transcribing
|
||||||
"""
|
"""
|
||||||
raw_segments: list[dict]
|
raw_segments: list[dict]
|
||||||
raw_segments = self._model.transcribe(path, **kwargs)["segments"] # type: ignore
|
raw_segments = self._model.transcribe(path, **kwargs)["segments"] # type: ignore
|
||||||
result: list[TimelineEvent] = []
|
result = AsrRawResult(
|
||||||
seg_id: int = 0
|
engine="whisper",
|
||||||
|
segments=[]
|
||||||
|
)
|
||||||
for raw_segment in raw_segments:
|
for raw_segment in raw_segments:
|
||||||
ev = TimelineEvent(
|
e = AsrRawSegment(
|
||||||
id = f"asr_{seg_id}",
|
start=float(raw_segment["start"]),
|
||||||
timestamp=float(raw_segment["start"]),
|
end=float(raw_segment["end"]),
|
||||||
duration=float(raw_segment["end"]) - float(raw_segment["start"]),
|
text=str(raw_segment["text"]),
|
||||||
payload=raw_segment["text"],
|
engine={
|
||||||
custom={
|
|
||||||
"whisper_temperature": float(raw_segment["temperature"]),
|
"whisper_temperature": float(raw_segment["temperature"]),
|
||||||
"whisper_avg_logprob": float(raw_segment["avg_logprob"]),
|
"whisper_avg_logprob": float(raw_segment["avg_logprob"]),
|
||||||
"whisper_no_speech_prob": float(raw_segment["no_speech_prob"])
|
"whisper_no_speech_prob": float(raw_segment["no_speech_prob"])
|
||||||
},
|
}
|
||||||
links=[]
|
|
||||||
)
|
)
|
||||||
result.append(ev)
|
result.segments.append(e)
|
||||||
seg_id += 1
|
|
||||||
return result
|
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 time
|
||||||
import os
|
import os
|
||||||
from dataclasses import asdict
|
from dataclasses import asdict
|
||||||
|
from enum import Enum
|
||||||
|
from typing import Callable
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
from asr import Asr, AsrRawResult
|
||||||
|
from asr_filter import AsrFilter
|
||||||
|
|
||||||
from transcriber import Transcriber
|
|
||||||
from agent import Agent
|
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:
|
def check_cuda() -> None:
|
||||||
"""Checks if CUDA is available."""
|
"""Checks if CUDA is available."""
|
||||||
if not torch.cuda.is_available():
|
if not torch.cuda.is_available():
|
||||||
@@ -28,33 +46,22 @@ def setup_arguments() -> argparse.Namespace:
|
|||||||
Returns:
|
Returns:
|
||||||
- argparse namespace
|
- argparse namespace
|
||||||
"""
|
"""
|
||||||
voice_models = Transcriber.get_models_list()
|
voice_models = Asr.get_models_list()
|
||||||
|
|
||||||
parser = argparse.ArgumentParser(
|
parser = argparse.ArgumentParser(
|
||||||
prog="sumka",
|
prog="sumka",
|
||||||
description="Summarizes large video/audio files into convenient format",
|
description="Summarizes large video/audio files into convenient format",
|
||||||
)
|
)
|
||||||
parser.add_argument("filename", type=str)
|
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--voice-model",
|
"--asr-model",
|
||||||
choices=voice_models,
|
choices=voice_models,
|
||||||
default="turbo" if "turbo" in voice_models else voice_models[-1]
|
default="turbo" if "turbo" in voice_models else voice_models[-1]
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--voice-language",
|
"--asr-language",
|
||||||
choices=["ru", "en"],
|
choices=["ru", "en"],
|
||||||
default="ru"
|
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(
|
parser.add_argument(
|
||||||
"--ai-model",
|
"--ai-model",
|
||||||
type=str,
|
type=str,
|
||||||
@@ -67,287 +74,124 @@ def setup_arguments() -> argparse.Namespace:
|
|||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--ai-api-key",
|
"--ai-api-key",
|
||||||
type=str,
|
type=str
|
||||||
default="gdsfgds"
|
|
||||||
)
|
)
|
||||||
parser.add_argument("-v", action='store_true')
|
parser.add_argument("-v", action='store_true')
|
||||||
return parser.parse_args()
|
return parser.parse_args()
|
||||||
|
|
||||||
#
|
#
|
||||||
# GENERIC
|
# Workflow
|
||||||
#
|
#
|
||||||
def make_windows(events: list[TimelineEvent], context_symbols: int, payload_symbols: int) -> list[AnalysisWindow]:
|
def on_media_separation(current_step: Step, input_data: dict | None) -> tuple[Step | None, dict | None]:
|
||||||
result: list[AnalysisWindow] = []
|
# there may be no data for this step
|
||||||
window_start = 0
|
if input_data:
|
||||||
while window_start < len(events):
|
raise RuntimeError("There must be no input data for Media Separation")
|
||||||
window = AnalysisWindow([], [], [])
|
# do not split if there's `audio.mp3`
|
||||||
result.append(window)
|
if os.path.isfile(WORKFLOW_DATA[current_step][0]):
|
||||||
# build the window itself
|
logging.info("Skipping media separation")
|
||||||
window_end = window_start + 1
|
return (Step.VOICE_RECOGNITION, None)
|
||||||
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 execute_timeline_request(res: TimelineProcessingResult, request: dict):
|
# filesnames to look for
|
||||||
req = request["req"]
|
VIDEO_INPUTS = [
|
||||||
if req == "modify":
|
"input.mp4",
|
||||||
id = request["id"]
|
"input.mkv",
|
||||||
payload = request["payload"]
|
"input.avi"
|
||||||
valid = [e for e in res.events if e.id == id]
|
]
|
||||||
if not len(valid):
|
AUDIO_INPUTS = [
|
||||||
raise RuntimeError(f"AI tries to modify nonexistent event with ID `{id}`")
|
"input.mp3",
|
||||||
valid[0].payload = payload
|
"input.m4a",
|
||||||
logging.info(f"Updated `{id}`'s payload to `{payload}`")
|
"input.wav"
|
||||||
else:
|
]
|
||||||
print(json.dumps(request, indent=2, ensure_ascii=False))
|
# 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 build_document_structure(timeline: TimelineProcessingResult, args: argparse.Namespace):
|
def on_voice_recognition(current_step: Step, input_data: dict | None) -> tuple[Step | None, dict | None]:
|
||||||
results_path = "document_structure.json"
|
# do not perform recognition if output file exists
|
||||||
result = []
|
if os.path.isfile(WORKFLOW_DATA[current_step][0]):
|
||||||
# create the agent
|
logging.info("Skipping voice recognition")
|
||||||
agent = Agent(
|
with open(WORKFLOW_DATA[current_step][0], "rb") as f:
|
||||||
model=args.ai_model,
|
return (Step.ASR_FILTER, json.load(f))
|
||||||
base_url=args.ai_base_url,
|
logging.info(f"Loading ASR model `{ARGS.asr_model}`, language `{ARGS.asr_language}`")
|
||||||
api_key=args.ai_api_key
|
asr = Asr(ARGS.asr_model)
|
||||||
)
|
logging.info(f"Speech recognition...")
|
||||||
# create the system prompt
|
result = asr.recognize("audio.mp3", language=ARGS.asr_language)
|
||||||
with open("prompts/build_structure.md", "r") as f:
|
logging.info(f"Speech recognition done")
|
||||||
system_prompt = AgentMessage(
|
return (Step.ASR_FILTER, result.model_dump(mode="json"))
|
||||||
f.read(),
|
|
||||||
"system"
|
def on_asr_filter(current_step: Step, input_data: dict | None) -> tuple[Step | None, dict | None]:
|
||||||
)
|
# do not filter if output file exists
|
||||||
windows = make_windows(timeline.events, args.window_context_size * 2, args.window_payload_size * 2)
|
if os.path.isfile(WORKFLOW_DATA[current_step][0]):
|
||||||
# context
|
logging.info("Skipping ASR filter")
|
||||||
context = {}
|
with open(WORKFLOW_DATA[current_step][0], "rb") as f:
|
||||||
# process each window
|
return (Step.ASR_EVENTS, json.load(f))
|
||||||
for window_id, window in enumerate(windows):
|
# bad request
|
||||||
# prepare request body
|
if input_data is None:
|
||||||
req = {
|
logging.error("Can'f filter raw ASR ouput without input_data")
|
||||||
"before_readonly": [e.get_ai_dict() for e in window.before_readonly],
|
return (None, None)
|
||||||
"content": [e.get_ai_dict() for e in window.modifiable],
|
logging.info("Filtering raw ASR output...")
|
||||||
"after_readonly": [e.get_ai_dict() for e in window.after_readonly],
|
filter = AsrFilter()
|
||||||
"document_context": context
|
result = filter.filter(AsrRawResult(**input_data))
|
||||||
}
|
return (Step.ASR_EVENTS, result.model_dump(mode="json"))
|
||||||
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
|
|
||||||
|
|
||||||
#
|
#
|
||||||
# AUDIO
|
# Main
|
||||||
#
|
#
|
||||||
def transcribe_audio(audio_path: str,
|
WORKFLOW_DATA: dict[Step, tuple[str, Callable[[Step, dict | None], tuple[Step | None, dict | None]] | None]] = {
|
||||||
args: argparse.Namespace) -> list[TimelineEvent]:
|
Step.MEDIA_SEPARATION: ("audio.mp3", on_media_separation),
|
||||||
"""Transcribes audio.
|
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.
|
||||||
|
|
||||||
Args:
|
Step.CODE: (
|
||||||
- audio_path - path to the audio file
|
"path/to/result.json",
|
||||||
- result_path - path to the resulting JSON file
|
(cur_step: Step, input_data: dict | None)
|
||||||
- args - arguments as returned by argsparse
|
-> (next_step: Step | None, output_data: dict | None)
|
||||||
|
)
|
||||||
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
|
|
||||||
|
|
||||||
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
|
|
||||||
|
|
||||||
def main() -> None:
|
def main() -> None:
|
||||||
"""Application entry point"""
|
"""Application entry point"""
|
||||||
|
global ARGS
|
||||||
check_cuda()
|
check_cuda()
|
||||||
args = setup_arguments()
|
ARGS = setup_arguments()
|
||||||
# setup the logger
|
logging.basicConfig(level=logging.DEBUG if ARGS.v else logging.INFO)
|
||||||
logging.basicConfig(level=logging.DEBUG if args.v else logging.INFO)
|
step_to_do = Step.MEDIA_SEPARATION
|
||||||
# transcribe
|
intermediate_result: dict | None = None
|
||||||
audio_events = transcribe_audio(
|
# execute steps while possible
|
||||||
args.filename,
|
while step_to_do:
|
||||||
args
|
step_data = WORKFLOW_DATA[step_to_do]
|
||||||
)
|
output_file_path = step_data[0]
|
||||||
# prepare audio windows
|
func = step_data[1]
|
||||||
audio_windows = prepare_audio_windows(
|
if not func:
|
||||||
audio_events,
|
logging.error(f"Step {step_to_do} has no function, stopping")
|
||||||
args
|
break
|
||||||
)
|
logging.info(f"Executing step {step_to_do}")
|
||||||
# process audio window
|
step_to_do, ret = func(step_to_do, intermediate_result)
|
||||||
processing_result = process_audio_windows(
|
if ret is not None:
|
||||||
audio_windows,
|
intermediate_result = ret
|
||||||
args
|
with open(output_file_path, "w") as f:
|
||||||
)
|
json.dump(ret, f, ensure_ascii=False, indent=4)
|
||||||
# build the document
|
logging.info(f"Intermediate results are saved {output_file_path}")
|
||||||
build_document_structure(processing_result, args)
|
# final report
|
||||||
|
logging.info(f"Workflow was interrupted at step {step_to_do}")
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
try:
|
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 dataclasses import dataclass
|
||||||
from typing import Literal
|
from typing import Literal, Any
|
||||||
|
import subprocess
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class TimelineEvent:
|
class Event:
|
||||||
id: str
|
"""Event within the timeline"""
|
||||||
"""ID of the event in format `asr_18`"""
|
|
||||||
|
id: int
|
||||||
|
"""Unique ID of the event (within single timeline)"""
|
||||||
|
|
||||||
|
type: Literal["voice", "vis", "ocr"]
|
||||||
|
"""Type of the event"""
|
||||||
|
|
||||||
timestamp: float
|
timestamp: float
|
||||||
"""When did the event happen"""
|
"""When did the event start (or happen, if `duration` is 0)"""
|
||||||
|
|
||||||
duration: float
|
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
|
text: str
|
||||||
"""Payload of the event (text for `asr`, description for `vis`, OCR result for `ocr`)"""
|
"""Textual representation of the event (or data in `payload`, if it is not None)"""
|
||||||
|
|
||||||
custom: dict
|
payload: str | None
|
||||||
"""Custom data"""
|
"""Payload of the event. `None` for `voice`, image path for `vis`, OCR result for `ocr`"""
|
||||||
|
|
||||||
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
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class AnalysisWindow:
|
class Timeline:
|
||||||
before_readonly: list[TimelineEvent]
|
"""Timeline of events"""
|
||||||
"""Timeline events for context (before current window)"""
|
|
||||||
|
|
||||||
modifiable: list[TimelineEvent]
|
events: list[Event]
|
||||||
"""Timeline events that are to be analyzed"""
|
"""Events of the timeline"""
|
||||||
|
|
||||||
after_readonly: list[TimelineEvent]
|
def ffmpeg_split_video(input_path: str) -> tuple[str, str]:
|
||||||
"""Timeline events for context (after current window)"""
|
"""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
|
# Audio
|
||||||
class AgentMessage:
|
"-map", "0:a:0",
|
||||||
content: str
|
"-vn",
|
||||||
"""Content of the message"""
|
"-c:a", "libmp3lame",
|
||||||
|
"-q:a", "2",
|
||||||
|
audio_path,
|
||||||
|
|
||||||
role: Literal["system", "assistant", "user"]
|
# Video
|
||||||
"""Who sent the message"""
|
"-map", "0:v:0",
|
||||||
|
"-an",
|
||||||
|
"-c:v", "libx264",
|
||||||
|
"-crf", "18",
|
||||||
|
"-preset", "fast",
|
||||||
|
video_path,
|
||||||
|
],
|
||||||
|
check=True,
|
||||||
|
)
|
||||||
|
return audio_path, video_path
|
||||||
|
|
||||||
@dataclass
|
def ffmpeg_to_mp3(input_path: str) -> str:
|
||||||
class TimelineProcessingResult:
|
"""Convert `input_path` path to mp3 and save it at ./audio.mp3"""
|
||||||
events: list[TimelineEvent]
|
output_path = "audio.mp3"
|
||||||
"""List of events"""
|
subprocess.run(
|
||||||
|
[
|
||||||
desired_events: list[TimelineEvent]
|
"ffmpeg",
|
||||||
"""List of events desired for existance"""
|
"-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