diff --git a/backend/schemas.py b/backend/schemas.py index 054c0c5..4349fbc 100644 --- a/backend/schemas.py +++ b/backend/schemas.py @@ -30,6 +30,10 @@ class JobStage(str, Enum): UPLOADING = "uploading" PREPROCESSING = "preprocessing" TRANSCRIBING = "transcribing" + # Distinct from ALIGNING: entered only while an HF-backed alignment model is + # being downloaded mid-transcription (not pre-provisioned), so the frozen + # "Aligning… 50%" bar reads as a download instead of a hang (#145). + DOWNLOADING_ALIGN_MODEL = "downloading_align_model" ALIGNING = "aligning" DIARIZING = "diarizing" EMOTION_ANALYSIS = "emotion_analysis" diff --git a/backend/services/align_models.py b/backend/services/align_models.py index ffbe218..e7617fb 100644 --- a/backend/services/align_models.py +++ b/backend/services/align_models.py @@ -13,8 +13,11 @@ from __future__ import annotations +import logging from collections.abc import Iterable +logger = logging.getLogger(__name__) + # Torch-native alignment languages: bundled by torchaudio into the shared # ``~/.cache/torch`` (not HuggingFace), so they are never part of provisioning's # HF pre-fetch and the transcriber lets WhisperX resolve them (``model_name=None``). @@ -79,3 +82,33 @@ def align_repos_for(languages: Iterable[str]) -> list[str]: if repo and repo not in repos: repos.append(repo) return repos + + +def align_model_cached(repo_id: str) -> bool: + """Whether ``repo_id``'s alignment model is already in the local HF cache. + + A ``config.json`` hit is the presence proxy — WhisperX loads the model via + ``Wav2Vec2ForCTC.from_pretrained``, which always needs it. Used to decide + whether loading the model will trigger a mid-transcription download that + should be surfaced as a distinct job stage rather than a frozen "Aligning…" + (#145). + + ``huggingface_hub`` is imported lazily (kept out of this module's top-level + imports; mirrors ``provisioning._snapshot_downloader``) so align resolution + stays usable in the lightweight CI env without the ML stack. The probe uses + the default cache dir, which ``huggingface_hub`` derives from ``HF_HOME`` at + import time — the same env the bundled service sets before first import — so + the probe and the subsequent load share cache resolution by construction. + + Degrades safe: any failure (including ``huggingface_hub`` being absent) + returns ``True`` ("assume present"), which merely omits the download + indicator and never blocks alignment — the ``ALIGNMENT_TIMEOUT_SEC`` watchdog + remains the safety net. + """ + try: + from huggingface_hub import try_to_load_from_cache + + return isinstance(try_to_load_from_cache(repo_id, "config.json"), str) + except Exception: # noqa: BLE001 - safe degrade; keep the load path unaffected + logger.debug("align cache probe failed for %s; assuming present", repo_id, exc_info=True) + return True diff --git a/backend/services/transcriber.py b/backend/services/transcriber.py index c02bfa1..d90ac4d 100644 --- a/backend/services/transcriber.py +++ b/backend/services/transcriber.py @@ -6,6 +6,7 @@ import threading import time from collections import defaultdict +from collections.abc import Callable from pathlib import Path import config @@ -18,7 +19,7 @@ MeetingStatus, TranscriptSegment, ) -from backend.services.align_models import HF_ALIGN_REPOS +from backend.services.align_models import HF_ALIGN_REPOS, align_model_cached from backend.services.job_queue import job_queue from backend.services.multilingual_transcriber import transcribe_multilingual from config import MEETINGS_DIR, WHISPER_BATCH_SIZE, WHISPER_DEVICE, WHISPER_MODEL @@ -60,6 +61,33 @@ def _worker() -> None: return "ok", box.get("value") +def _load_align_model_watchdogged( + job_id: str, + load_fn: Callable[[], object], + *, + align_model_name: str | None, + align_progress: int, +) -> tuple[str, object]: + """Run the alignment-model load under the watchdog, surfacing a download stage. + + ``load_fn`` is a zero-arg callable closed over the caller's lazily-imported + ``whisperx`` (keeps this helper ML-import-free). When the model is HF-backed + (``align_model_name`` set) and not already cached, the load will download it + mid-transcription — reported as a distinct ``downloading_align_model`` stage so + the 50%/82% bar reads as a download rather than a frozen "Aligning…" (#145) — + then reset to ``aligning`` once loaded. Torch-native models (``align_model_name`` + is ``None``, shared ``~/.cache/torch``) never flip the stage. Returns + ``_call_with_timeout``'s ``(status, loaded)``; the watchdog remains the safety net. + """ + downloading = bool(align_model_name) and not align_model_cached(align_model_name) + if downloading: + job_queue.update_job(job_id, stage="downloading_align_model", progress=align_progress) + status, loaded = _call_with_timeout(load_fn, ALIGNMENT_TIMEOUT_SEC, "align-load") + if downloading: + job_queue.update_job(job_id, stage="aligning", progress=align_progress) + return status, loaded + + # Languages with a wav2vec2 alignment model available in WhisperX (single-language path). ALIGNMENT_LANGUAGES = { "en", @@ -413,13 +441,15 @@ def _update_transcribe_progress(): align_model_name = HF_ALIGN_REPOS.get(detected_language) # Load (and, on first use, download) the alignment model under a watchdog # so a stalled fetch degrades to segment-level timestamps instead of - # hanging at the 50% stage (debug session macos-transcribe-stuck-70). - status, loaded = _call_with_timeout( + # hanging at the 50% stage (debug session macos-transcribe-stuck-70). A + # not-yet-cached fetch surfaces a distinct download stage (#145). + status, loaded = _load_align_model_watchdogged( + job_id, lambda: whisperx.load_align_model( language_code=detected_language, device=device, model_name=align_model_name ), - ALIGNMENT_TIMEOUT_SEC, - "align-load", + align_model_name=align_model_name, + align_progress=50, ) if status != "ok": logger.error( @@ -504,7 +534,7 @@ def _finalize_aligned(segments: list[dict], group: list[dict], language: str) -> return finalized -def _align_multilingual_segments(ml_segments: list[dict], audio, device: str) -> list[dict]: +def _align_multilingual_segments(job_id: str, ml_segments: list[dict], audio, device: str) -> list[dict]: """Group segments by detected language and align each group with its own model. Segments are grouped by their ``language`` field and each group is aligned with @@ -530,10 +560,14 @@ def _align_multilingual_segments(ml_segments: list[dict], audio, device: str) -> align_model_name = HF_ALIGN_REPOS.get(language) # Watchdog the load/download so a stalled fetch degrades to # segment-level timestamps instead of hanging (see single-language path). - status, loaded = _call_with_timeout( - lambda: whisperx.load_align_model(language_code=language, device=device, model_name=align_model_name), - ALIGNMENT_TIMEOUT_SEC, - "align-load", + # A not-yet-cached fetch surfaces a distinct download stage (#145). + status, loaded = _load_align_model_watchdogged( + job_id, + lambda language=language, align_model_name=align_model_name: whisperx.load_align_model( + language_code=language, device=device, model_name=align_model_name + ), + align_model_name=align_model_name, + align_progress=82, ) if status == "timeout": raise TimeoutError(f"alignment model load exceeded {ALIGNMENT_TIMEOUT_SEC}s") @@ -601,7 +635,7 @@ def _progress(pct: int) -> None: # Per-language word alignment (BR-8, BR-9, EC-3). job_queue.update_job(job_id, stage="aligning", progress=82) - aligned = _align_multilingual_segments(ml_segments, audio, device) + aligned = _align_multilingual_segments(job_id, ml_segments, audio, device) # Diarize + assign speakers (shared with the single-language path). result, diarize_turns = _diarize_and_assign(job_id, audio, num_speakers, device, {"segments": aligned}, progress=85) diff --git a/docs/plans/145-align-model-download-indicator.md b/docs/plans/145-align-model-download-indicator.md new file mode 100644 index 0000000..e3f4df9 --- /dev/null +++ b/docs/plans/145-align-model-download-indicator.md @@ -0,0 +1,65 @@ +# Plan: Surface a model-download indicator when an alignment model is fetched mid-transcription + +**Story**: #145 +**Spec**: N/A (follow-up to #141; no spec document) +**Branch**: feature/145-align-model-download-indicator +**Date**: 2026-09-04 +**Mode**: Standard — mocking-heavy ML pipeline; tests written alongside each change. + +## Technical Decisions + +### TD-1: Distinct job stage `downloading_align_model` +- **Context**: A lazy align-model download inside the align stage leaves the job frozen at `stage=aligning, progress=50`, indistinguishable from a hang. +- **Decision**: Add `DOWNLOADING_ALIGN_MODEL = "downloading_align_model"` to `JobStage` and set it while the model is being fetched (issue option 1). +- **Alternatives considered**: Reuse provisioning's `DownloadState` (rejected — separate lifecycle, provisioning phase vs job stage); gate/announce the download up front (rejected — the detected language is unknown until transcription runs). + +### TD-2: Cache-presence proxy via `huggingface_hub.try_to_load_from_cache` +- **Context**: We must decide whether the align model will download before entering the load call. +- **Decision**: New `align_model_cached(repo_id)` in `align_models.py` probes `try_to_load_from_cache(repo_id, "config.json")`; True iff a cached path string is returned. huggingface_hub imported lazily inside the function (preserves the module's import-light contract). Any exception (incl. `ModuleNotFoundError` when hf_hub is absent in the lightweight CI env) → return True = "assume present, skip the indicator" = today's behavior. Logged at debug so the degrade is observable. +- **Alternatives considered**: Full `scan_cache_dir` (overkill); checking model weights file (filename varies: `pytorch_model.bin` vs `model.safetensors`). `config.json` is a sufficient proxy since whisperx loads via `Wav2Vec2ForCTC.from_pretrained`, which always needs it. The `ALIGNMENT_TIMEOUT_SEC` watchdog covers a false "present". + +### TD-3: Only HF-backed models get the indicator +- **Context**: Torch-native align models (en/fr/de/es/it, `model_name=None`) live in the shared `~/.cache/torch` and are unaffected (per the issue). +- **Decision**: Skip the probe entirely when `align_model_name is None`. + +## Files to Create or Modify + +- `backend/schemas.py` — add `DOWNLOADING_ALIGN_MODEL` to `JobStage`. +- `backend/services/align_models.py` — add `align_model_cached(repo_id)` (lazy hf_hub import; debug-logged safe degrade). +- `backend/services/transcriber.py` — import `align_model_cached`; add `_load_align_model_watchdogged(...)` that sets the download stage when the HF-backed model is not cached, runs the watchdogged load, resets the stage to `aligning`, and returns `(status, loaded)`. The caller passes the whisperx-closed load callable so the helper stays ML-import-free. Wire both align sites (single-language progress=50, multilingual progress=82). +- `frontend/js/components/transcript-viewer.js` — add `downloading_align_model: 'Downloading alignment model...'` to `stageLabels`. +- `macos/Sources/MeetingTranscriberKit/Presentation/JobPresentation.swift` — add `"downloading_align_model": "Downloading alignment model"`. +- `tests/unit/test_align_models.py` — `align_model_cached`: cached→True, not-cached (None)→False, known-missing sentinel→False, hf_hub absent/raises→True. +- `tests/unit/test_transcriber.py` — HF-backed (Thai) run with probe False → `update_job` receives a `stage="downloading_align_model"` call; probe True → no such call; English (`model_name=None`) → no such call. Captured via a recording mock on `update_job` (queue keeps only latest state). +- `macos/Sources/MeetingTranscriberKitTests/UploadValidationTests.swift` — `label(for: "downloading_align_model") == "Downloading alignment model"`. + +## Approach per AC + +### AC: When a non-pre-provisioned HF align model must download, the job reports a distinct downloading state +Probe cache before the load; set `stage=downloading_align_model` (progress unchanged) while `whisperx.load_align_model` fetches; reset to `aligning` once loaded. Both align sites route through `_load_align_model_watchdogged`. + +### AC: Torch-native and already-cached models keep the current behavior +Probe skipped when `model_name is None`; when cached, no stage flip. + +### AC: The downloading state is surfaced in both UIs +Web `stageLabels` and Swift `JobStagePresentation` map the new stage to "Downloading alignment model". + +### AC: The watchdog remains the safety net +`_call_with_timeout(..., ALIGNMENT_TIMEOUT_SEC, "align-load")` is unchanged; each site's divergent outcome handling (single-language logs-and-degrades; multilingual raises) is preserved. + +## Commit Sequence + +1. `[#145]` schema enum + `align_models.align_model_cached` + unit tests +2. `[#145]` transcriber helper wiring both align sites + transcriber tests +3. `[#145]` frontend + swift download-stage labels + swift test + +## Risks and Trade-offs + +- `config.json` presence is a heuristic: cached config but missing weights → false "present" (no indicator), but the watchdog still covers it. +- Very fast downloads → brief flash of the downloading label. Acceptable. +- The probe uses the default HF cache_dir, which huggingface_hub derives from `HF_HOME` at import time — the same env the bundled service sets before first import (mirrors #141 provisioning), so probe and load share cache resolution by construction. Recorded as a code comment. + +## Deviations from Plan + +- `_align_multilingual_segments` gained a leading `job_id: str` parameter (not called out in the plan) so it can pass the job id into `_load_align_model_watchdogged`. Both the caller in `_run_multilingual_transcription` and the two direct test call sites were updated accordingly. +- No other deviations. The helper takes the whisperx-closed load callable (architect Nit) and the debug-logged safe degrade + shared-cache comment (architect Minors) as approved. diff --git a/frontend/js/components/transcript-viewer.js b/frontend/js/components/transcript-viewer.js index afbf946..2df8c0e 100644 --- a/frontend/js/components/transcript-viewer.js +++ b/frontend/js/components/transcript-viewer.js @@ -588,6 +588,7 @@ async function updateProgress(meetingId, jobId) { uploading: 'Uploading...', preprocessing: 'Preprocessing audio...', transcribing: 'Transcribing audio...', + downloading_align_model: 'Downloading alignment model...', aligning: 'Aligning timestamps...', diarizing: 'Identifying speakers...', saving: 'Saving results...', diff --git a/macos/Sources/MeetingTranscriberKit/Presentation/JobPresentation.swift b/macos/Sources/MeetingTranscriberKit/Presentation/JobPresentation.swift index 20b6886..e78ae6b 100644 --- a/macos/Sources/MeetingTranscriberKit/Presentation/JobPresentation.swift +++ b/macos/Sources/MeetingTranscriberKit/Presentation/JobPresentation.swift @@ -7,6 +7,7 @@ public enum JobStagePresentation { "uploading": "Uploading", "preprocessing": "Preparing audio", "transcribing": "Transcribing", + "downloading_align_model": "Downloading alignment model", "aligning": "Aligning timestamps", "diarizing": "Identifying speakers", "emotion_analysis": "Analyzing emotion", diff --git a/macos/Sources/MeetingTranscriberKitTests/UploadValidationTests.swift b/macos/Sources/MeetingTranscriberKitTests/UploadValidationTests.swift index 371a394..4bf085b 100644 --- a/macos/Sources/MeetingTranscriberKitTests/UploadValidationTests.swift +++ b/macos/Sources/MeetingTranscriberKitTests/UploadValidationTests.swift @@ -29,6 +29,9 @@ func runUploadValidationTests() { suite("JobStagePresentation.label") { expectEqual(JobStagePresentation.label(for: "transcribing"), "Transcribing", "known stage") expectEqual(JobStagePresentation.label(for: "diarizing"), "Identifying speakers", "friendly label") + expectEqual( + JobStagePresentation.label(for: "downloading_align_model"), + "Downloading alignment model", "align-model download stage (#145)") expectEqual(JobStagePresentation.label(for: ""), "Working", "empty → generic") expectEqual(JobStagePresentation.label(for: "custom_stage"), "Custom Stage", "unknown → humanized") } diff --git a/tests/unit/test_align_models.py b/tests/unit/test_align_models.py index 14d4002..1303657 100644 --- a/tests/unit/test_align_models.py +++ b/tests/unit/test_align_models.py @@ -1,5 +1,8 @@ from __future__ import annotations +import sys +import types + import pytest from backend.services import align_models @@ -32,6 +35,51 @@ def test_empty_input(self): assert align_models.align_repos_for([]) == [] +class TestAlignModelCached: + """align_model_cached probes the HF cache to decide if a load will download.""" + + def _install_fake_hub(self, monkeypatch, return_value): + """Inject a fake huggingface_hub whose try_to_load_from_cache is stubbed.""" + fake = types.ModuleType("huggingface_hub") + calls: list[tuple] = [] + + def _probe(repo_id, filename): + calls.append((repo_id, filename)) + if isinstance(return_value, Exception): + raise return_value + return return_value + + fake.try_to_load_from_cache = _probe + monkeypatch.setitem(sys.modules, "huggingface_hub", fake) + return calls + + def test_cached_path_string_is_present(self, monkeypatch): + calls = self._install_fake_hub(monkeypatch, "/cache/models--foo/config.json") + assert align_models.align_model_cached("foo/bar") is True + assert calls == [("foo/bar", "config.json")] + + def test_not_cached_none_is_absent(self, monkeypatch): + self._install_fake_hub(monkeypatch, None) + assert align_models.align_model_cached("foo/bar") is False + + def test_known_missing_sentinel_is_absent(self, monkeypatch): + # huggingface_hub returns a _CACHED_NO_EXIST sentinel (not a str) for a + # file known to be absent from the repo. + sentinel = object() + self._install_fake_hub(monkeypatch, sentinel) + assert align_models.align_model_cached("foo/bar") is False + + def test_probe_error_degrades_to_present(self, monkeypatch): + self._install_fake_hub(monkeypatch, RuntimeError("cache scan failed")) + assert align_models.align_model_cached("foo/bar") is True + + def test_missing_huggingface_hub_degrades_to_present(self, monkeypatch): + # Absent in the lightweight CI env: the lazy import raises, and the probe + # must degrade to "assume present" rather than blow up. + monkeypatch.setitem(sys.modules, "huggingface_hub", None) + assert align_models.align_model_cached("foo/bar") is True + + class TestHfAlignReposDriftGuard: """Guard that the copied HF map stays in sync with the installed WhisperX. diff --git a/tests/unit/test_transcriber.py b/tests/unit/test_transcriber.py index ce95451..33cdfdf 100644 --- a/tests/unit/test_transcriber.py +++ b/tests/unit/test_transcriber.py @@ -1114,7 +1114,7 @@ class TestAlignMultilingualSegments: def _run(ml_segments, align_side_effect): mw = _fake_whisperx(align_side_effect) with patch.dict("sys.modules", {"whisperx": mw, "torch": MagicMock()}): - result = _align_multilingual_segments(ml_segments, MagicMock(), "cpu") + result = _align_multilingual_segments("job-1", ml_segments, MagicMock(), "cpu") return result, mw def test_each_language_aligned_with_its_own_model(self): @@ -1233,7 +1233,7 @@ def _blocking_load(*a, **k): ml = [{"start": 1.0, "end": 3.0, "text": "Bonjour", "language": "fr"}] with patch.dict("sys.modules", {"whisperx": mw, "torch": MagicMock()}): - result = _align_multilingual_segments(ml, MagicMock(), "cpu") + result = _align_multilingual_segments("job-1", ml, MagicMock(), "cpu") release.set() seg = result[0] @@ -1243,6 +1243,59 @@ def _blocking_load(*a, **k): assert "words" not in seg # never aligned; word data absent +class TestLoadAlignModelWatchdogged: + """The align-model load surfaces a distinct download stage when an HF-backed + model is not yet cached, so a mid-transcription fetch is not mistaken for a + hang (#145). The queue keeps only the latest state, so stages are recorded + via a spy on update_job.""" + + @staticmethod + def _record_stages(queue: JobQueue, job_id: str) -> list[str]: + stages: list[str] = [] + original = queue.update_job + + def _spy(jid, **kwargs): + if kwargs.get("stage") is not None: + stages.append(kwargs["stage"]) + return original(jid, **kwargs) + + queue.update_job = _spy # type: ignore[method-assign] + return stages + + def _run(self, *, align_model_name, cached): + queue = JobQueue() + job = queue.create_job("m1") + stages = self._record_stages(queue, job.id) + with ( + patch.object(transcriber_module, "job_queue", queue), + patch.object(transcriber_module, "align_model_cached", return_value=cached) as probe, + ): + status, loaded = transcriber_module._load_align_model_watchdogged( + job.id, + lambda: ("model", "metadata"), + align_model_name=align_model_name, + align_progress=50, + ) + assert status == "ok" + assert loaded == ("model", "metadata") + return stages, probe + + def test_hf_backed_uncached_surfaces_download_then_aligning(self): + stages, probe = self._run(align_model_name="jonatasgrosman/wav2vec2-x", cached=False) + assert stages == ["downloading_align_model", "aligning"] + probe.assert_called_once_with("jonatasgrosman/wav2vec2-x") + + def test_hf_backed_cached_does_not_flip_stage(self): + stages, _probe = self._run(align_model_name="jonatasgrosman/wav2vec2-x", cached=True) + assert "downloading_align_model" not in stages + + def test_torch_native_skips_probe_and_stage(self): + # model_name=None (en/fr/de/es/it): shared torch cache, never a download stage. + stages, probe = self._run(align_model_name=None, cached=False) + assert "downloading_align_model" not in stages + probe.assert_not_called() + + class TestDiarizationWatchdog: """The diarization step must degrade to no-diarization instead of hanging forever when the pipeline stalls (debug session macos-transcribe-stuck-70)."""