Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion src/asr/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -238,7 +238,7 @@ def _run_transcription(
) -> int:
run_id = f"run-{uuid.uuid4().hex[:8]}"
collector = MetricsCollectorObserver() if verbose else None
observers = [ConsoleProgressObserver()]
observers = [ConsoleProgressObserver(verbose=verbose)]
if collector is not None:
observers.append(collector)
observer = ObserverMux(
Expand Down
69 changes: 51 additions & 18 deletions src/asr/observability/console.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,22 @@
from asr.observability.events import ObservabilityEvent

_PROGRESS_BAR_WIDTH = 20
_QUIET_STEP_LABELS = {
"preflight": "checking environment",
"prepare": "preparing audio",
"preprocess_vad": "detecting speech",
"transcribe": "transcribing",
}
_QUIET_FINALIZING_STEPS = frozenset(
{
"provider_merge",
"render_srt",
"render_vtt",
"render_json",
"write_outputs",
}
)
_STATUS_LABELS = {"ok": "done"}


class _ThinProgressBarColumn(ProgressColumn):
Expand All @@ -37,6 +53,7 @@ class ConsoleProgressObserver:
stream: TextIO = field(default_factory=lambda: sys.stdout)
warning_stream: TextIO = field(default_factory=lambda: sys.stderr)
is_tty: bool | None = None
verbose: bool = False
_current_index: int = field(default=0, init=False, repr=False)
_current_total: int = field(default=0, init=False, repr=False)
_current_name: str = field(default="", init=False, repr=False)
Expand All @@ -48,6 +65,7 @@ class ConsoleProgressObserver:
_console: Console | None = field(default=None, init=False, repr=False)
_progress: Progress | None = field(default=None, init=False, repr=False)
_progress_task_id: int | None = field(default=None, init=False, repr=False)
_quiet_phase: str | None = field(default=None, init=False, repr=False)

def __post_init__(self) -> None:
if self.is_tty is None:
Expand All @@ -62,15 +80,21 @@ def on_event(self, event: ObservabilityEvent) -> None:
self._current_window_index = 0
self._current_window_count = 0
self._file_start_perf = event.perf_counter
self._quiet_phase = None
self._write_line(self._with_elapsed("discover", event.perf_counter))
return

if event.event_type == "step_start":
if event.step == "provider_window":
self._quiet_phase = None
self._record_window_progress(event)
self._write_window_progress(event.perf_counter)
return
if not self._should_display_step(event):
self._handle_hidden_step(event)
return
self._stop_progress()
self._quiet_phase = None
self._write_line(self._with_elapsed(self._display_step(event), event.perf_counter))
return

Expand All @@ -80,30 +104,41 @@ def on_event(self, event: ObservabilityEvent) -> None:

if event.event_type == "file_end":
status = str(event.meta.get("status", "ok"))
if self._current_window_count > 0:
if status == "ok":
self._current_window_index = self._current_window_count
else:
self._stop_progress()
self._write_line(
f"{status} | {self._window_progress_line(event.perf_counter)}",
finalize=True,
)
return
self._write_window_progress(event.perf_counter)
self._stop_progress()
return
self._stop_progress()
self._write_line(
self._with_elapsed(status, event.perf_counter),
self._with_elapsed(self._display_status(status), event.perf_counter),
finalize=True,
)

def close(self) -> None:
self._stop_progress()

def _should_display_step(self, event: ObservabilityEvent) -> bool:
if self.verbose:
return True
return event.step in _QUIET_STEP_LABELS

def _handle_hidden_step(self, event: ObservabilityEvent) -> None:
if (
not self.verbose
and event.step in _QUIET_FINALIZING_STEPS
and self._progress is not None
):
self._stop_progress()
if self._quiet_phase != "finalizing":
self._quiet_phase = "finalizing"
self._write_line(self._with_elapsed("finalizing", event.perf_counter))
return
self._stop_progress()

def _display_step(self, event: ObservabilityEvent) -> str:
return event.step or "step"
step = event.step or "step"
if self.verbose:
return step
return _QUIET_STEP_LABELS.get(step, step)

def _display_status(self, status: str) -> str:
return _STATUS_LABELS.get(status, status)

def _record_window_progress(self, event: ObservabilityEvent) -> None:
window_index = event.meta.get("window_index")
Expand Down Expand Up @@ -180,8 +215,6 @@ def _stop_progress(self) -> None:
if self._progress is None:
return
self._progress.stop()
if self.is_tty:
self._last_width = 0
self._progress = None
self._progress_task_id = None

Expand Down Expand Up @@ -225,8 +258,8 @@ def _write_line(self, step: str, *, finalize: bool = False) -> None:
if self.is_tty:
tail = "\n" if finalize else ""
padded = line.ljust(self._last_width)
self._last_width = max(self._last_width, len(line))
self.stream.write(f"\r{padded}{tail}")
self._last_width = 0 if finalize else max(self._last_width, len(line))
else:
self.stream.write(line + "\n")
self.stream.flush()
Expand Down
32 changes: 29 additions & 3 deletions src/asr/providers/qwen_mlx.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
from __future__ import annotations

import math
import os
import shutil
import subprocess
import tempfile
Expand Down Expand Up @@ -115,9 +116,7 @@ def transcribe(
audio_path: Path,
speech_plan: SpeechPlan | None = None,
) -> TranscriptionDocument:
load = self._load_backend()
self._asr_model = self._asr_model or load(self.asr_model_id)
self._aligner_model = self._aligner_model or load(self.aligner_model_id)
self._ensure_models_loaded()

with observe_step(
self._observer,
Expand Down Expand Up @@ -580,6 +579,33 @@ def _load_backend(self):
) from exc
return load

def _ensure_models_loaded(self) -> None:
if self._asr_model is not None and self._aligner_model is not None:
return

self._suppress_model_download_progress()
load = self._load_backend()
if self._asr_model is None:
self._asr_model = load(self.asr_model_id)
if self._aligner_model is None:
self._aligner_model = load(self.aligner_model_id)

def _suppress_model_download_progress(self) -> None:
flag = os.environ.get("HF_HUB_DISABLE_PROGRESS_BARS")
if (
flag is not None
and flag.strip().lower() in {"0", "false", "off", "no"}
):
return

os.environ.setdefault("HF_HUB_DISABLE_PROGRESS_BARS", "1")

try:
from huggingface_hub.utils import disable_progress_bars
except ImportError:
return
disable_progress_bars()
Comment on lines +593 to +607

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

The logic for suppressing Hugging Face download progress bars is effective, but it relies on modifying the global os.environ. While acceptable for a CLI tool, consider using os.environ.setdefault or only setting it if it's not already present to minimize side effects. Additionally, the check for explicitly_enabled could be simplified.

Suggested change
def _suppress_model_download_progress(self) -> None:
flag = os.environ.get("HF_HUB_DISABLE_PROGRESS_BARS")
if flag is None:
os.environ["HF_HUB_DISABLE_PROGRESS_BARS"] = "1"
explicitly_enabled = False
else:
explicitly_enabled = flag.strip().lower() in {"0", "false", "off", "no"}
if explicitly_enabled:
return
try:
from huggingface_hub.utils import disable_progress_bars
except ImportError:
return
disable_progress_bars()
def _suppress_model_download_progress(self) -> None:
flag = os.environ.get("HF_HUB_DISABLE_PROGRESS_BARS")
if flag is not None and flag.strip().lower() in {"0", "false", "off", "no"}:
return
os.environ.setdefault("HF_HUB_DISABLE_PROGRESS_BARS", "1")
try:
from huggingface_hub.utils import disable_progress_bars
disable_progress_bars()
except ImportError:
pass


def _probe_duration_sec(self, audio_path: Path) -> float:
return probe_duration_sec(audio_path)

Expand Down
29 changes: 29 additions & 0 deletions tests/test_cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -281,6 +281,35 @@ def test_main_enables_console_progress_by_default(

self.assertEqual(exit_code, 0)
self.assertTrue(mock_console_observer.called)
self.assertFalse(mock_console_observer.call_args.kwargs["verbose"])

@patch("asr.cli.ConsoleProgressObserver")
@patch("asr.cli.discover_cli_sources")
@patch("asr.cli.run_environment_preflight")
@patch("asr.cli.process_media_file")
def test_main_passes_verbose_to_console_progress(
self,
mock_process,
mock_preflight,
mock_discover,
mock_console_observer,
) -> None:
with TemporaryDirectory() as tmp:
source = Path(tmp) / "demo.mov"
source.write_text("x", encoding="utf-8")
output_root = Path(tmp) / "outputs"
mock_discover.return_value = [(source, Path(tmp))]
mock_preflight.return_value = (True, "")
mock_process.return_value = TranscriptionDocument(
source_path=str(source.with_suffix(".wav")),
provider_name="fake",
segments=[],
)

exit_code = main([str(source), "--verbose", "--output-dir", str(output_root)])

self.assertEqual(exit_code, 0)
self.assertTrue(mock_console_observer.call_args.kwargs["verbose"])

@patch("asr.cli.discover_cli_sources")
@patch("asr.cli.run_environment_preflight")
Expand Down
Loading