Huge refactoring

This commit is contained in:
2026-09-17 00:46:04 +03:00
parent c65a9a91e2
commit 018be076b9
11 changed files with 534 additions and 344 deletions

406
main.py
View File

@@ -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: