Initial commit
This commit is contained in:
358
main.py
Normal file
358
main.py
Normal file
@@ -0,0 +1,358 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user