Retain live processing timings across UI reruns
This commit is contained in:
+167
-49
@@ -2,7 +2,11 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from datetime import date
|
||||
from queue import Empty, Queue
|
||||
from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
@@ -35,6 +39,7 @@ from mka.application.people_yaml import (
|
||||
import_people_yaml,
|
||||
)
|
||||
from mka.application.performance import DEFAULT_PERFORMANCE_PROFILE, PERFORMANCE_PROFILES
|
||||
from mka.application.progress_timing import TimingSnapshot
|
||||
from mka.application.run_inputs import (
|
||||
RunInputJsonError,
|
||||
RunInputState,
|
||||
@@ -338,40 +343,132 @@ def _render_participants() -> list[ParticipantInput]:
|
||||
return people
|
||||
|
||||
|
||||
def _progress_callback(
|
||||
def _format_duration(seconds: float) -> str:
|
||||
"""Format numeric seconds as MM:SS or H:MM:SS."""
|
||||
whole_seconds = max(0, int(seconds))
|
||||
hours, remainder = divmod(whole_seconds, 3600)
|
||||
minutes, seconds = divmod(remainder, 60)
|
||||
return f"{hours}:{minutes:02d}:{seconds:02d}" if hours else f"{minutes:02d}:{seconds:02d}"
|
||||
|
||||
|
||||
def _render_progress_status(
|
||||
status_box: Any,
|
||||
stage_table: Any,
|
||||
progress_slot: Any,
|
||||
states: dict[str, str],
|
||||
) -> Any:
|
||||
progress_bar = None
|
||||
timing: TimingSnapshot,
|
||||
message: str,
|
||||
) -> None:
|
||||
suffix = " …" if timing.running else ""
|
||||
status_box.info(f"{message} — total runtime {_format_duration(timing.total_duration)}{suffix}")
|
||||
stage_table.table(
|
||||
[
|
||||
{
|
||||
"Stage": STAGE_LABELS[stage],
|
||||
"Status": states[stage],
|
||||
"Runtime": (
|
||||
_format_duration(timing.stage_durations[stage])
|
||||
+ (" …" if timing.active_stage == stage else "")
|
||||
if stage in timing.stage_durations
|
||||
else ""
|
||||
),
|
||||
}
|
||||
for stage in STAGES
|
||||
]
|
||||
+ [
|
||||
{
|
||||
"Stage": "Total runtime",
|
||||
"Status": "",
|
||||
"Runtime": _format_duration(timing.total_duration) + suffix,
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
def update(event: AppProgressEvent) -> None:
|
||||
nonlocal progress_bar
|
||||
if event.stage in states:
|
||||
states[event.stage] = "running" if event.status == "started" else event.status
|
||||
if event.stage == "failed":
|
||||
running = next(
|
||||
|
||||
def _apply_progress_event(event: AppProgressEvent, states: dict[str, str]) -> None:
|
||||
if event.stage in states:
|
||||
states[event.stage] = "running" if event.status == "started" else event.status
|
||||
if event.stage == "failed":
|
||||
running = next((stage for stage, status in states.items() if status == "running"), None)
|
||||
if running:
|
||||
states[running] = "failed"
|
||||
|
||||
|
||||
def _run_with_live_progress(
|
||||
service: MeetingProcessingService,
|
||||
work: Callable[[Callable[[AppProgressEvent], None]], Any],
|
||||
states: dict[str, str],
|
||||
initial_message: str,
|
||||
) -> tuple[Any, dict[str, str], str, TimingSnapshot]:
|
||||
"""Run pipeline work while rendering the shared timer and progress events."""
|
||||
status_box = st.empty()
|
||||
stage_table = st.empty()
|
||||
progress_slot = st.empty()
|
||||
event_queue: Queue[AppProgressEvent] = Queue()
|
||||
latest_message = initial_message
|
||||
progress_bar = None
|
||||
with ThreadPoolExecutor(max_workers=1) as executor:
|
||||
future = executor.submit(work, event_queue.put)
|
||||
while not future.done():
|
||||
try:
|
||||
while True:
|
||||
event = event_queue.get_nowait()
|
||||
_apply_progress_event(event, states)
|
||||
latest_message = event.message or STAGE_LABELS.get(event.stage, event.stage)
|
||||
if event.progress is not None:
|
||||
if progress_bar is None:
|
||||
progress_bar = progress_slot.progress(0.0)
|
||||
progress_bar.progress(
|
||||
min(max(event.progress, 0.0), 1.0),
|
||||
text=(
|
||||
f"{STAGE_LABELS.get(event.stage, event.stage)}: "
|
||||
f"{event.progress:.0%}"
|
||||
),
|
||||
)
|
||||
except Empty:
|
||||
pass
|
||||
_render_progress_status(
|
||||
status_box,
|
||||
stage_table,
|
||||
states,
|
||||
service.timer.snapshot(),
|
||||
latest_message,
|
||||
)
|
||||
time.sleep(0.2)
|
||||
try:
|
||||
result = future.result()
|
||||
except Exception:
|
||||
while not event_queue.empty():
|
||||
event = event_queue.get_nowait()
|
||||
_apply_progress_event(event, states)
|
||||
latest_message = event.message or STAGE_LABELS.get(event.stage, event.stage)
|
||||
service.timer.finish_run()
|
||||
running_stage = next(
|
||||
(stage for stage, status in states.items() if status == "running"),
|
||||
None,
|
||||
)
|
||||
if running:
|
||||
states[running] = "failed"
|
||||
elapsed = f"{event.elapsed_seconds:.1f} s"
|
||||
message = event.message or STAGE_LABELS.get(event.stage, event.stage)
|
||||
status_box.info(f"{message} — elapsed {elapsed}")
|
||||
stage_table.table(
|
||||
[{"Stage": STAGE_LABELS[stage], "Status": states[stage]} for stage in STAGES]
|
||||
)
|
||||
if event.progress is not None:
|
||||
if progress_bar is None:
|
||||
progress_bar = progress_slot.progress(0.0)
|
||||
progress_bar.progress(
|
||||
min(max(event.progress, 0.0), 1.0),
|
||||
text=f"{STAGE_LABELS.get(event.stage, event.stage)}: {event.progress:.0%}",
|
||||
)
|
||||
if running_stage is not None:
|
||||
states[running_stage] = "failed"
|
||||
snapshot = service.timer.snapshot()
|
||||
_render_progress_status(status_box, stage_table, states, snapshot, latest_message)
|
||||
st.session_state["processing_progress"] = (states.copy(), latest_message, snapshot)
|
||||
raise
|
||||
while not event_queue.empty():
|
||||
event = event_queue.get_nowait()
|
||||
_apply_progress_event(event, states)
|
||||
latest_message = event.message or STAGE_LABELS.get(event.stage, event.stage)
|
||||
running_stage = next((stage for stage, status in states.items() if status == "running"), None)
|
||||
if running_stage is not None:
|
||||
states[running_stage] = "completed" if getattr(result, "succeeded", True) else "failed"
|
||||
snapshot = service.timer.snapshot()
|
||||
_render_progress_status(status_box, stage_table, states, snapshot, latest_message)
|
||||
st.session_state["processing_progress"] = (states.copy(), latest_message, snapshot)
|
||||
return result, states, latest_message, snapshot
|
||||
|
||||
return update
|
||||
|
||||
def _remember_regeneration_timing(
|
||||
states: dict[str, str], message: str, timing: TimingSnapshot
|
||||
) -> None:
|
||||
st.session_state["regeneration_progress"] = (states.copy(), message, timing)
|
||||
|
||||
|
||||
def _speaker_options(
|
||||
@@ -447,6 +544,9 @@ def _render_result() -> None:
|
||||
return
|
||||
st.divider()
|
||||
st.header("Protocol result")
|
||||
if retained_progress := st.session_state.get("processing_progress"):
|
||||
states, message, timing = retained_progress
|
||||
_render_progress_status(st.empty(), st.empty(), states, timing, message)
|
||||
if not outcome.succeeded:
|
||||
st.error(f"Processing failed during {outcome.failed_stage}: {outcome.error_message}")
|
||||
if outcome.run_dir:
|
||||
@@ -494,18 +594,33 @@ def _render_result() -> None:
|
||||
disabled=duplicate_assignments,
|
||||
type="primary",
|
||||
):
|
||||
st.session_state.pop("regeneration_progress", None)
|
||||
run_dir = outcome.run_dir
|
||||
mapping_items = tuple(selections.items())
|
||||
selected_profile = st.session_state.get("performance_profile", DEFAULT_PERFORMANCE_PROFILE)
|
||||
states = {stage: "skipped" for stage in STAGES}
|
||||
states["protocol_generation"] = "pending"
|
||||
try:
|
||||
with st.spinner("Regenerating protocol without rerunning audio processing..."):
|
||||
regenerated = service.regenerate_protocol(
|
||||
outcome.run_dir,
|
||||
selections,
|
||||
performance_profile=st.session_state.get(
|
||||
"performance_profile", DEFAULT_PERFORMANCE_PROFILE
|
||||
),
|
||||
)
|
||||
regenerated, states, message, timing = _run_with_live_progress(
|
||||
service,
|
||||
lambda progress_sink: service.regenerate_protocol(
|
||||
run_dir,
|
||||
dict(mapping_items),
|
||||
progress_sink=progress_sink,
|
||||
performance_profile=selected_profile,
|
||||
),
|
||||
states,
|
||||
"Starting protocol regeneration",
|
||||
)
|
||||
except (OSError, RuntimeError, ValueError) as exc:
|
||||
_remember_regeneration_timing(
|
||||
states,
|
||||
"Protocol regeneration failed",
|
||||
service.timer.snapshot(),
|
||||
)
|
||||
st.error(f"Protocol regeneration failed: {exc}")
|
||||
else:
|
||||
_remember_regeneration_timing(states, message, timing)
|
||||
st.session_state.outcome = regenerated
|
||||
_queue_edited_protocol(regenerated.original_protocol or "")
|
||||
st.session_state.speaker_mapping_message = (
|
||||
@@ -618,31 +733,34 @@ def main() -> None:
|
||||
)
|
||||
meeting_id = stable_id(title)
|
||||
audio_path = service.preserve_upload(meeting_id, audio.name, audio)
|
||||
st.session_state.pop("regeneration_progress", None)
|
||||
st.session_state.pop("processing_progress", None)
|
||||
st.header("Processing status")
|
||||
status_box = st.empty()
|
||||
stage_table = st.empty()
|
||||
progress_slot = st.empty()
|
||||
states = {stage: "pending" for stage in STAGES}
|
||||
if not diarization_enabled:
|
||||
states["diarization"] = "skipped"
|
||||
callback = _progress_callback(status_box, stage_table, progress_slot, states)
|
||||
outcome = service.process(
|
||||
audio_path,
|
||||
meeting,
|
||||
participants,
|
||||
ProcessingOptions(
|
||||
diarization_enabled=diarization_enabled,
|
||||
audio_normalization=audio_normalization,
|
||||
performance_profile=performance_profile,
|
||||
outcome, _, _, _ = _run_with_live_progress(
|
||||
service,
|
||||
lambda progress_sink: service.process(
|
||||
audio_path,
|
||||
meeting,
|
||||
participants,
|
||||
ProcessingOptions(
|
||||
diarization_enabled=diarization_enabled,
|
||||
audio_normalization=audio_normalization,
|
||||
performance_profile=performance_profile,
|
||||
),
|
||||
progress_sink=progress_sink,
|
||||
),
|
||||
progress_sink=callback,
|
||||
states,
|
||||
"Starting processing",
|
||||
)
|
||||
st.session_state.outcome = outcome
|
||||
st.session_state.edited_protocol = outcome.original_protocol or ""
|
||||
if outcome.succeeded:
|
||||
status_box.success("Processing completed.")
|
||||
st.success("Processing completed.")
|
||||
else:
|
||||
status_box.error(
|
||||
st.error(
|
||||
f"Processing failed during {outcome.failed_stage}: "
|
||||
f"{outcome.error_message} Artifacts were preserved."
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user