Huge refactoring
This commit is contained in:
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:
|
||||
|
||||
Reference in New Issue
Block a user