358 lines
12 KiB
Python
358 lines
12 KiB
Python
"""Application entry point"""
|
|
|
|
import traceback
|
|
import argparse
|
|
import logging
|
|
import json
|
|
import time
|
|
import os
|
|
from dataclasses import asdict
|
|
|
|
import torch
|
|
|
|
from transcriber import Transcriber
|
|
from agent import Agent
|
|
from utils import TimelineEvent, AnalysisWindow, AgentMessage, TimelineProcessingResult
|
|
|
|
def check_cuda() -> None:
|
|
"""Checks if CUDA is available."""
|
|
if not torch.cuda.is_available():
|
|
raise RuntimeError(
|
|
"CUDA is unavaiable! Refusing to start, it's pointless."
|
|
)
|
|
|
|
def setup_arguments() -> argparse.Namespace:
|
|
"""Parses CLI arguments and returns namespace. It may terminate the
|
|
application on invalid arguments.
|
|
|
|
Returns:
|
|
- argparse namespace
|
|
"""
|
|
voice_models = Transcriber.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",
|
|
choices=voice_models,
|
|
default="turbo" if "turbo" in voice_models else voice_models[-1]
|
|
)
|
|
parser.add_argument(
|
|
"--voice-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,
|
|
default="GigaChat-3-Pro"
|
|
)
|
|
parser.add_argument(
|
|
"--ai-base-url",
|
|
type=str,
|
|
default="https://api.giga.chat/v1"
|
|
)
|
|
parser.add_argument(
|
|
"--ai-api-key",
|
|
type=str,
|
|
default="gdsfgds"
|
|
)
|
|
parser.add_argument("-v", action='store_true')
|
|
return parser.parse_args()
|
|
|
|
#
|
|
# GENERIC
|
|
#
|
|
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 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 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
|
|
|
|
#
|
|
# AUDIO
|
|
#
|
|
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
|
|
|
|
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:
|
|
"""Application entry point"""
|
|
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)
|
|
|
|
if __name__ == "__main__":
|
|
try:
|
|
main()
|
|
except SystemExit:
|
|
raise
|
|
except:
|
|
traceback.print_exc() |