Skip to content

Commit 7e5f81b

Browse files
authored
Merge pull request #24 from MapleEve/fix/model-load-timing
chore: log model load timings
2 parents e4c7170 + 93ace27 commit 7e5f81b

13 files changed

Lines changed: 588 additions & 45 deletions

File tree

app/application/transcription_jobs.py

Lines changed: 9 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
"""Application-level transcription job orchestration."""
22

33
import logging
4+
import time
45
from pathlib import Path
56

67
from config import (
@@ -39,6 +40,7 @@ def _record_status(status: str) -> None:
3940
extra_filename = audio_path.name if status == "converting" else None
4041
_write_status(job_id, status, filename=extra_filename)
4142

43+
job_started = time.perf_counter()
4244
try:
4345

4446
def _process_pipeline():
@@ -69,16 +71,20 @@ def _process_pipeline():
6971
jobs[job_id]["result"] = tr
7072
_write_status(job_id, "completed")
7173
logger.info(
72-
"Job %s completed: %d segments, %d speakers",
73-
job_id,
74+
"transcription_job_timing status=completed elapsed_s=%.3f segment_count=%d speaker_count=%d",
75+
time.perf_counter() - job_started,
7476
len(tr.get("segments", [])),
7577
len(tr.get("speaker_map", {})),
7678
)
7779
if file_hash:
7880
unregister_in_flight(file_hash, job_id)
7981

8082
except Exception as e:
81-
logger.exception("Job %s failed", job_id)
83+
logger.exception(
84+
"transcription_job_timing status=failed elapsed_s=%.3f error_type=%s",
85+
time.perf_counter() - job_started,
86+
e.__class__.__name__,
87+
)
8288
jobs[job_id]["status"] = "failed"
8389
jobs[job_id]["error"] = str(e)
8490
_write_status(job_id, "failed", error=str(e))

app/pipeline/orchestrator.py

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@
1414
import json
1515
import re
1616
import tempfile
17+
import time
1718
from contextlib import nullcontext
1819
from pathlib import Path
1920
from typing import Any
@@ -336,11 +337,23 @@ def whisper(self):
336337
whisper_device,
337338
compute_type,
338339
)
340+
load_started = time.perf_counter()
339341
self._whisper = WhisperModel(
340342
model_ref,
341343
**_faster_whisper_device_kwargs(whisper_device),
342344
compute_type=compute_type,
343345
)
346+
logger.info(
347+
"Loaded faster-whisper model in %.2fs (cold_load=True, device=%s, compute_type=%s)",
348+
time.perf_counter() - load_started,
349+
whisper_device,
350+
compute_type,
351+
)
352+
else:
353+
logger.info(
354+
"Reusing faster-whisper model (hot reuse, device=%s)",
355+
getattr(self, "_whisper_device", None) or getattr(self, "device", ""),
356+
)
344357
return self._whisper
345358

346359
@property
@@ -363,6 +376,7 @@ def diarization(self):
363376
token=self.hf_token,
364377
)
365378
logger.info("Loading pyannote diarization model")
379+
load_started = time.perf_counter()
366380
self._diarization = _load_trusted_pyannote_model(
367381
PyannotePipeline.from_pretrained,
368382
model_ref,
@@ -371,6 +385,11 @@ def diarization(self):
371385
_dev = diarization_device if ":" in diarization_device else "cuda:0"
372386
if diarization_device.startswith("cuda"):
373387
self._diarization.to(torch.device(_dev))
388+
logger.info(
389+
"Loaded pyannote diarization model in %.2fs (cold_load=True, device=%s)",
390+
time.perf_counter() - load_started,
391+
diarization_device,
392+
)
374393
# Suppress over-segmentation of short backchannel turns
375394
try:
376395
if hasattr(self._diarization, "_binarize") and hasattr(
@@ -385,6 +404,12 @@ def diarization(self):
385404
)
386405
except Exception as exc:
387406
logger.warning("Could not set min_duration_off: %s", exc)
407+
else:
408+
logger.info(
409+
"Reusing pyannote diarization model (hot reuse, device=%s)",
410+
getattr(self, "_diarization_device", None)
411+
or getattr(self, "device", ""),
412+
)
388413
return self._diarization
389414

390415
@property
@@ -400,6 +425,7 @@ def embedding_model(self):
400425
)
401426
model_ref = _resolve_local_pyannote_file(model_ref, "pytorch_model.bin")
402427
logger.info("Loading WeSpeaker speaker encoder")
428+
load_started = time.perf_counter()
403429
model = _load_trusted_pyannote_model(
404430
Model.from_pretrained,
405431
model_ref,
@@ -409,6 +435,16 @@ def embedding_model(self):
409435
# window="whole" returns one embedding vector per full chunk —
410436
# exactly what we need for per-turn embeddings.
411437
self._embedding_model = Inference(model, window="whole")
438+
logger.info(
439+
"Loaded WeSpeaker speaker encoder in %.2fs (cold_load=True, device=%s)",
440+
time.perf_counter() - load_started,
441+
embedding_device,
442+
)
443+
else:
444+
logger.info(
445+
"Reusing WeSpeaker speaker encoder (hot reuse, device=%s)",
446+
getattr(self, "_embedding_device", None) or getattr(self, "device", ""),
447+
)
412448
return self._embedding_model
413449

414450
def transcribe(

app/pipeline/runner.py

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
from __future__ import annotations
44

55
import logging
6+
import time
67
from typing import Any
78

89
from infra.audio import cleanup_generated_files
@@ -15,6 +16,35 @@
1516
DEFAULT_STAGE_ORDER = available_stage_slots()
1617

1718

19+
def _safe_stage_metrics(context: PipelineContext, stage_name: str) -> dict[str, Any]:
20+
metrics: dict[str, Any] = {}
21+
stage_metadata = context.metadata.get(stage_name)
22+
if isinstance(stage_metadata, dict):
23+
for key in (
24+
"status",
25+
"model",
26+
"language",
27+
"segment_count",
28+
"speaker_count",
29+
"turn_count",
30+
"applied",
31+
"reason",
32+
"persisted",
33+
):
34+
if key in stage_metadata:
35+
metrics[key] = stage_metadata[key]
36+
if "segment_count" not in metrics and context.aligned_segments:
37+
metrics["segment_count"] = len(context.aligned_segments)
38+
if "speaker_count" not in metrics:
39+
if context.speaker_embeddings:
40+
metrics["speaker_count"] = len(context.speaker_embeddings)
41+
elif context.diarization_turns:
42+
metrics["speaker_count"] = len(
43+
{turn.get("speaker") for turn in context.diarization_turns}
44+
)
45+
return metrics
46+
47+
1848
class PipelineRunner:
1949
"""Execute the stable stage order against the current pipeline implementation."""
2050

@@ -42,7 +72,21 @@ def run_context(self, pipeline: Any, request: PipelineRequest) -> PipelineContex
4272
context.metadata.setdefault("selected_providers", {})[stage_name] = (
4373
request.provider_for(stage_name)
4474
)
75+
stage_started = time.perf_counter()
4576
stage(context)
77+
elapsed_s = time.perf_counter() - stage_started
78+
metrics = _safe_stage_metrics(context, stage_name)
79+
context.metadata.setdefault("stage_timings", {})[stage_name] = round(
80+
elapsed_s,
81+
3,
82+
)
83+
logger.info(
84+
"pipeline_stage_timing stage=%s elapsed_s=%.3f provider=%s metrics=%s",
85+
stage_name,
86+
elapsed_s,
87+
request.provider_for(stage_name),
88+
metrics,
89+
)
4690
return context
4791
finally:
4892
cleanup_generated_files(context.temporary_paths)

app/providers/asr/default.py

Lines changed: 14 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
from __future__ import annotations
44

55
import logging
6+
import time
67
from collections import Counter
78
from typing import Any
89

@@ -191,7 +192,9 @@ def run_faster_whisper_asr(
191192
if no_repeat_ngram_size and no_repeat_ngram_size >= 3:
192193
whisper_kwargs["no_repeat_ngram_size"] = no_repeat_ngram_size
193194

194-
segments_iter, info = pipeline.whisper.transcribe(audio_path, **whisper_kwargs)
195+
whisper_model = pipeline.whisper
196+
processing_started = time.perf_counter()
197+
segments_iter, info = whisper_model.transcribe(audio_path, **whisper_kwargs)
195198
raw_segments = [
196199
{
197200
"start": round(float(segment.start), 3),
@@ -200,8 +203,18 @@ def run_faster_whisper_asr(
200203
}
201204
for segment in segments_iter
202205
]
206+
processing_elapsed_s = time.perf_counter() - processing_started
203207
segments, hallucination_guard = suppress_repetition_hallucinations(raw_segments)
204208
detected_language = info.language
209+
audio_duration_s = max((segment["end"] for segment in raw_segments), default=0.0)
210+
logger.info(
211+
"asr_processing_timing model=faster-whisper elapsed_s=%.3f language=%s segment_count=%d raw_segment_count=%d duration_s=%.3f",
212+
processing_elapsed_s,
213+
detected_language,
214+
len(segments),
215+
len(raw_segments),
216+
audio_duration_s,
217+
)
205218
logger.info(
206219
"Transcription done: %d segments, language=%s, repetition_guard=%s",
207220
len(segments),

app/providers/diarization/default.py

Lines changed: 31 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55
import logging
66
import os
77
import re
8+
import time
89
from contextlib import contextmanager
910
from collections.abc import Callable
1011
from inspect import Parameter, signature
@@ -186,7 +187,9 @@ def run_pyannote_diarization(
186187
if max_speakers:
187188
kwargs["max_speakers"] = max_speakers
188189

189-
result = pipeline.diarization(audio_path, **kwargs)
190+
diarization_model = pipeline.diarization
191+
processing_started = time.perf_counter()
192+
result = diarization_model(audio_path, **kwargs)
190193
turns: list[dict[str, object]] = []
191194
for turn, _, speaker in result.itertracks(yield_label=True):
192195
turns.append(
@@ -196,6 +199,15 @@ def run_pyannote_diarization(
196199
"speaker": speaker,
197200
}
198201
)
202+
elapsed_s = time.perf_counter() - processing_started
203+
logger.info(
204+
"diarization_processing_timing model=pyannote elapsed_s=%.3f device=%s turn_count=%d speaker_count=%d",
205+
elapsed_s,
206+
getattr(pipeline, "_diarization_device", None)
207+
or getattr(pipeline, "device", ""),
208+
len(turns),
209+
len({turn["speaker"] for turn in turns}),
210+
)
199211
return turns
200212

201213

@@ -242,10 +254,20 @@ def align_diarized_segments_with_metadata(
242254
language,
243255
pipeline.device,
244256
)
257+
load_started = time.perf_counter()
245258
with _cache_only_alignment_environment():
246259
align_model, align_metadata = whisperx.load_align_model(
247260
**load_kwargs,
248261
)
262+
logger.info(
263+
"Loaded WhisperX alignment model in %.2fs "
264+
"(cold_load=True, language=%s, model_source=%s, device=%s)",
265+
time.perf_counter() - load_started,
266+
language,
267+
model_source,
268+
pipeline.device,
269+
)
270+
processing_started = time.perf_counter()
249271
aligned_result = whisperx.align(
250272
segments,
251273
align_model,
@@ -254,7 +276,15 @@ def align_diarized_segments_with_metadata(
254276
pipeline.device,
255277
return_char_alignments=False,
256278
)
279+
processing_elapsed_s = time.perf_counter() - processing_started
257280
segments = aligned_result.get("segments", segments)
281+
logger.info(
282+
"alignment_processing_timing model=whisperx elapsed_s=%.3f language=%s segment_count=%d device=%s",
283+
processing_elapsed_s,
284+
language,
285+
len(segments),
286+
pipeline.device,
287+
)
258288
logger.info("WhisperX forced alignment succeeded for language=%s", language)
259289
metadata = {
260290
"status": "succeeded",

app/providers/embedding/default.py

Lines changed: 14 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
from __future__ import annotations
44

55
import logging
6+
import time
67

78
import numpy as np
89
import torchaudio
@@ -50,8 +51,7 @@ def extract_embeddings_for_turns(
5051
)
5152
except Exception as exc:
5253
logger.warning(
53-
"Failed to load segment %s [%d:%d]: %s",
54-
speaker,
54+
"Failed to load embedding audio segment [%d:%d]: %s",
5555
start_sample,
5656
end_sample,
5757
exc,
@@ -66,18 +66,30 @@ def extract_embeddings_for_turns(
6666
speaker_segments.setdefault(speaker, []).append(chunk)
6767

6868
embeddings: dict[str, np.ndarray] = {}
69+
model_processing_elapsed_s = 0.0
70+
processed_chunk_count = 0
6971
for speaker, chunks in speaker_segments.items():
7072
emb_list = []
7173
chunks.sort(key=lambda chunk: chunk.shape[1], reverse=True)
7274
for chunk in chunks[:10]:
7375
embedding_model = pipeline.embedding_model
7476
embedding_device = getattr(pipeline, "embedding_device", pipeline.device)
77+
processing_started = time.perf_counter()
7578
emb = embedding_model(
7679
{"waveform": chunk.to(embedding_device), "sample_rate": target_sr}
7780
)
81+
model_processing_elapsed_s += time.perf_counter() - processing_started
82+
processed_chunk_count += 1
7883
emb_list.append(np.asarray(emb))
7984
if emb_list:
8085
embeddings[speaker] = np.mean(emb_list, axis=0)
86+
logger.info(
87+
"embedding_processing_timing model=wespeaker elapsed_s=%.3f device=%s speaker_count=%d chunk_count=%d",
88+
model_processing_elapsed_s,
89+
getattr(pipeline, "embedding_device", getattr(pipeline, "device", "")),
90+
len(embeddings),
91+
processed_chunk_count,
92+
)
8193
return embeddings
8294

8395

0 commit comments

Comments
 (0)