From d253aa2b0ea2562ed4e1702380fac7e22bd11828 Mon Sep 17 00:00:00 2001 From: Vuong Hoang Date: Wed, 30 Sep 2026 15:00:46 -0700 Subject: [PATCH] feat(parakeet): selectable model and a retry for speech that got no words Parakeet v2/v3 sometimes stop emitting for tens of seconds deep inside a long full-attention chunk while someone is talking; the same audio transcribed on its own is usually fine. After stitching, any stretch of at least --retry-gaps seconds (default 3) with no word, where most 25 ms frames sit within 12 dB of the recording's typical speech level, is re-transcribed on its own (in pieces of at most 60 s) and the words that land inside it are spliced in; segments are then rebuilt from the words with the model's own rule (end after . ? !). Output that triggers no retry is unchanged. The same retry runs in the short-audio script, which shares the helpers. The retry is best effort: a failing piece is skipped and a failing retry keeps the first pass. A piece that looks like a hallucination loop (mostly one repeated token, or more than 7 words a second) is not spliced in, and a retried word that repeats its neighbour at the gap edge or a piece boundary is dropped (the first pass's copy wins). Chunk and retry audio go through a per-run temporary directory instead of fixed /tmp paths. PARAKEET_MODEL_PATH (absolute, or relative to the env) selects the .nemo both scripts load; the default stays parakeet-tdt-0.6b-v3.nemo. The JSON "model" field reports the file actually loaded, and the Go adapter records it as ModelUsed instead of a hardcoded name. The CLI and JSON seam is otherwise unchanged: the new flag is optional, and the JSON gains retried_gaps. --- .../adapters/parakeet_adapter.go | 7 +- .../adapters/py/nvidia/parakeet_transcribe.py | 53 +++- .../py/nvidia/parakeet_transcribe_buffered.py | 231 +++++++++++++----- .../py/nvidia/tests/test_parakeet_slicing.py | 159 +++++++++++- 4 files changed, 380 insertions(+), 70 deletions(-) diff --git a/internal/transcription/adapters/parakeet_adapter.go b/internal/transcription/adapters/parakeet_adapter.go index 4fa252d..5f3bb9e 100644 --- a/internal/transcription/adapters/parakeet_adapter.go +++ b/internal/transcription/adapters/parakeet_adapter.go @@ -342,7 +342,9 @@ func (p *ParakeetAdapter) Transcribe(ctx context.Context, input interfaces.Audio } result.ProcessingTime = time.Since(startTime) - result.ModelUsed = "parakeet-tdt-0.6b-v3" + if result.ModelUsed == "" { + result.ModelUsed = "parakeet-tdt-0.6b-v3" + } result.Metadata = p.CreateDefaultMetadata(params) logger.Info("Parakeet transcription completed", @@ -525,6 +527,8 @@ func (p *ParakeetAdapter) parseResult(tempDir string, input interfaces.AudioInpu End float64 `json:"end"` } `json:"segment_timestamps"` Confidence interface{} `json:"confidence,omitempty"` + // The scripts report the model they actually loaded (PARAKEET_MODEL_PATH can change it). + Model string `json:"model"` } if err := json.Unmarshal(data, ¶keetResult); err != nil { @@ -538,6 +542,7 @@ func (p *ParakeetAdapter) parseResult(tempDir string, input interfaces.AudioInpu Segments: make([]interfaces.TranscriptSegment, len(parakeetResult.SegmentTimestamps)), WordSegments: make([]interfaces.TranscriptWord, len(parakeetResult.WordTimestamps)), Confidence: 0.0, // Default confidence + ModelUsed: parakeetResult.Model, } // Convert segments diff --git a/internal/transcription/adapters/py/nvidia/parakeet_transcribe.py b/internal/transcription/adapters/py/nvidia/parakeet_transcribe.py index 235e52a..6b2539e 100644 --- a/internal/transcription/adapters/py/nvidia/parakeet_transcribe.py +++ b/internal/transcription/adapters/py/nvidia/parakeet_transcribe.py @@ -7,9 +7,17 @@ import argparse import json import sys import os +import shutil +import tempfile from pathlib import Path +import librosa import nemo.collections.asr as nemo_asr +# Shared with the long-audio script, which lives in the same directory. +from parakeet_transcribe_buffered import ( + DEFAULT_RETRY_GAP_SECS, resolve_model_path, retry_speech_gaps, transcribe_chunk, +) + def transcribe_audio( audio_path: str, @@ -18,14 +26,11 @@ def transcribe_audio( context_left: int = 256, context_right: int = 256, include_confidence: bool = True, + retry_gap_secs: float = DEFAULT_RETRY_GAP_SECS, ): """ Transcribe audio using NVIDIA Parakeet model. """ - # Determine model path - model_filename = "parakeet-tdt-0.6b-v3.nemo" - model_path = None - # Locate project root: derived from VIRTUAL_ENV, which is set by `uv run` to path/.venv virtual_env = os.environ.get("VIRTUAL_ENV") if not virtual_env: @@ -33,10 +38,11 @@ def transcribe_audio( sys.exit(1) project_root = os.path.dirname(virtual_env) - model_path = os.path.join(project_root, model_filename) + model_path = resolve_model_path(project_root) + model_name = os.path.splitext(os.path.basename(model_path))[0] if not os.path.exists(model_path): - print(f"Error during transcription: Can't find {model_filename} in project root: {project_root}") + print(f"Error during transcription: Can't find model file: {model_path}") sys.exit(1) print(f"Loading NVIDIA Parakeet model from: {model_path}") @@ -80,6 +86,27 @@ def transcribe_audio( text = result_data.text word_timestamps = result_data.timestamp.get("word", []) segment_timestamps = result_data.timestamp.get("segment", []) + retried = 0 + words_added = False + if retry_gap_secs > 0: + workdir = tempfile.mkdtemp(prefix="parakeet-") # per run: concurrent jobs share /tmp + try: + audio, sr = librosa.load(audio_path, sr=None, mono=True) + + def transcribe_span(start_sample, end_sample): + return transcribe_chunk(asr_model, audio[start_sample:end_sample], sr, + start_sample / sr, os.path.join(workdir, "retry.wav")) + before = len(word_timestamps) + word_timestamps, segment_timestamps, retried = retry_speech_gaps( + word_timestamps, segment_timestamps, audio, sr, retry_gap_secs, transcribe_span + ) + words_added = len(word_timestamps) != before + if words_added: + text = " ".join(w["word"] for w in word_timestamps) + except Exception as err: # best effort: never lose a finished transcription + print(f"Warning: gap retry failed ({err}); keeping the first pass") + finally: + shutil.rmtree(workdir, ignore_errors=True) print(f"Transcription: {text}") @@ -90,14 +117,15 @@ def transcribe_audio( "word_timestamps": word_timestamps, "segment_timestamps": segment_timestamps, "audio_file": audio_path, - "model": "parakeet-tdt-0.6b-v3", + "model": model_name, + "retried_gaps": retried, "context": { "left": context_left, "right": context_right } } - if include_confidence: + if include_confidence and not words_added: # first-pass scores no longer line up # Add confidence scores if available if hasattr(result_data, 'confidence') and result_data.confidence: output_data["confidence"] = result_data.confidence @@ -119,7 +147,7 @@ def transcribe_audio( "transcription": text, "language": "en", "audio_file": audio_path, - "model": "parakeet-tdt-0.6b-v3" + "model": model_name } if output_file: @@ -154,6 +182,12 @@ def main(): "--context-right", type=int, default=256, help="Right attention context size (default: 256)" ) + parser.add_argument( + "--retry-gaps", type=float, default=DEFAULT_RETRY_GAP_SECS, + help=f"Re-transcribe on its own any stretch of at least this many seconds where " + f"the audio holds speech but the model produced no words " + f"(default: {DEFAULT_RETRY_GAP_SECS}; 0 disables)" + ) parser.add_argument( "--include-confidence", action="store_true", default=True, help="Include confidence scores" @@ -178,6 +212,7 @@ def main(): context_left=args.context_left, context_right=args.context_right, include_confidence=args.include_confidence, + retry_gap_secs=args.retry_gaps, ) except Exception as e: print(f"Error during transcription: {e}") diff --git a/internal/transcription/adapters/py/nvidia/parakeet_transcribe_buffered.py b/internal/transcription/adapters/py/nvidia/parakeet_transcribe_buffered.py index ba755c1..49bb72e 100644 --- a/internal/transcription/adapters/py/nvidia/parakeet_transcribe_buffered.py +++ b/internal/transcription/adapters/py/nvidia/parakeet_transcribe_buffered.py @@ -12,15 +12,21 @@ import argparse import json import sys import os +import shutil +import tempfile import librosa import soundfile as sf import numpy as np from pathlib import Path +DEFAULT_MODEL = "parakeet-tdt-0.6b-v3.nemo" DEFAULT_OVERLAP_SECS = 4.0 DEFAULT_PAUSE_SEARCH_SECS = 0.0 # opt-in; measured no gain on top of the overlap QUIET_WINDOW_SECS = 0.3 SAME_WORD_SECS = 0.5 +DEFAULT_RETRY_GAP_SECS = 3.0 +RETRY_PIECE_SECS = 60.0 +RETRY_MAX_WORDS_PER_SEC = 7.0 def plan_slices(audio, sr, max_chunk_secs, overlap_secs=0.0, search_secs=0.0): @@ -103,14 +109,7 @@ def stitch_slices(slice_results, cut_times, chunk_spans=None): elif len(kept) == stop - start: segments.append(seg) elif kept: - segments.append({ - **seg, - "segment": " ".join(w["word"] for w in kept), - "start_offset": kept[0]["start_offset"], - "end_offset": kept[-1]["end_offset"], - "start": kept[0]["start"], - "end": kept[-1]["end"], - }) + segments.append({**seg, **_segment_of(kept)}) return words, segments @@ -159,6 +158,134 @@ def _segment_ranges(words, segments): return ranges +def resolve_model_path(project_root): + """The .nemo to load: $PARAKEET_MODEL_PATH (absolute, or relative to the env), + else the bundled v3 weights.""" + chosen = os.environ.get("PARAKEET_MODEL_PATH", "").strip() or DEFAULT_MODEL + return chosen if os.path.isabs(chosen) else os.path.join(project_root, chosen) + + +def find_speech_gaps(words, audio, sr, min_gap): + """Stretches of at least `min_gap` s with no word, where at least half the + 25 ms frames are within 12 dB of the recording's typical speech level: the + model went quiet while someone was talking. Parakeet sometimes does this for + tens of seconds deep inside a long chunk; the same audio transcribed on its + own is usually fine.""" + hop, win = int(0.010 * sr), int(0.025 * sr) + frames = max(0, (len(audio) - win) // hop) + if frames == 0: + return [] + # A strided view, reduced in blocks: bounded memory even for hours of audio. + view = np.lib.stride_tricks.sliding_window_view(audio, win)[::hop][:frames] + power = np.empty(frames) + for k in range(0, frames, 8192): + power[k:k + 8192] = np.mean(view[k:k + 8192].astype(np.float64) ** 2, axis=1) + level = 10 * np.log10(power + 1e-12) + speech = float(np.median(level[level >= np.median(level)])) + duration = len(audio) / sr + gaps = [] + for a, b in zip([0.0] + [w["end"] for w in words], [w["start"] for w in words] + [duration]): + if b - a >= min_gap: + span = level[int(a * sr) // hop:int(b * sr) // hop] + if len(span) and np.mean(span >= speech - 12) >= 0.5: + gaps.append((a, b)) + return gaps + + +def retry_speech_gaps(words, segments, audio, sr, min_gap, transcribe): + """Re-transcribe each speech gap on its own, in pieces of at most + RETRY_PIECE_SECS, and splice in the words that land inside it. + + `transcribe(start_sample, end_sample)` returns (text, words, segments) in + absolute time. Returns (words, segments, number of gaps retried); when words + were added, segments are rebuilt from the words with the model's own rule. + """ + gaps = find_speech_gaps(words, audio, sr, min_gap) + if not gaps: + return words, segments, 0 + found = [] + for a, b in gaps: + t = a + while t < b: + e = min(b, t + RETRY_PIECE_SECS) + try: + _, piece, _ = transcribe(max(0, int((t - 0.5) * sr)), min(len(audio), int((e + 0.5) * sr))) + except Exception as err: # best effort: the first pass already succeeded + print(f"Warning: retry of {t:.1f}-{e:.1f}s failed ({err}); keeping the first pass there") + piece = [] + piece = [w for w in piece if t <= w["start"] < e] + if _plausible(piece, e - t): + found.extend(piece) + t = e + if not found: + return words, segments, len(gaps) + merged = _without_retried_duplicates(sorted(words + found, key=lambda w: w["start"]), + {id(w) for w in found}) + return merged, segments_from_words(merged), len(gaps) + + +def _plausible(piece, seconds): + """False for what looks like a hallucination loop rather than speech: many words + that are mostly one repeated token, or more words per second than anyone says.""" + if len(piece) >= 10 and len({_normalize(w["word"]) for w in piece}) < 0.3 * len(piece): + return False + return len(piece) <= RETRY_MAX_WORDS_PER_SEC * max(seconds, 1.0) + + +def _without_retried_duplicates(merged, retried): + """The retry's padding re-hears the words either side of a gap, and a word + near a piece boundary is heard by both pieces. Drop a retried word that + repeats its neighbour within SAME_WORD_SECS; the first pass's copy wins.""" + out = [] + for w in merged: + if out and _normalize(out[-1]["word"]) == _normalize(w["word"]) \ + and w["start"] - out[-1]["start"] <= SAME_WORD_SECS: + if id(w) in retried: + continue + if id(out[-1]) in retried: + out[-1] = w + continue + out.append(w) + return out + + +def segments_from_words(words): + """Segments as Parakeet's decoding config makes them: a segment ends after a + word ending in '.', '?' or '!' (segment_seperators, no gap threshold).""" + segments, current = [], [] + for w in words: + current.append(w) + if w["word"].endswith((".", "?", "!")): + segments.append(_segment_of(current)) + current = [] + if current: + segments.append(_segment_of(current)) + return segments + + +def _segment_of(words): + return {"segment": " ".join(w["word"] for w in words), + "start_offset": words[0]["start_offset"], "end_offset": words[-1]["end_offset"], + "start": words[0]["start"], "end": words[-1]["end"]} + + +def transcribe_chunk(asr_model, audio, sr, start_time, path): + """Transcribe one chunk via a WAV file; word and segment times become absolute.""" + sf.write(path, audio, sr) + try: + result = asr_model.transcribe([path], batch_size=1, timestamps=True)[0] + finally: + if os.path.exists(path): + os.remove(path) + words, segments = [], [] + if getattr(result, "timestamp", None): + words = [dict(w, start=w["start"] + start_time, end=w["end"] + start_time) + for w in result.timestamp.get("word", [])] + segments = [dict(g, start=g["start"] + start_time, end=g["end"] + start_time) + for g in result.timestamp.get("segment", [])] + return result.text, words, segments + + def split_audio_file(audio_path, chunk_duration_secs=300, overlap_secs=0.0, search_secs=0.0): """Split audio file into chunks of at most chunk_duration_secs.""" audio, sr = librosa.load(audio_path, sr=None, mono=True) @@ -173,7 +300,7 @@ def split_audio_file(audio_path, chunk_duration_secs=300, overlap_secs=0.0, sear 'duration': len(chunk_audio) / sr }) - return chunks, sr, [cut / sr for cut in cuts] + return chunks, sr, [cut / sr for cut in cuts], audio def transcribe_buffered( @@ -182,16 +309,13 @@ def transcribe_buffered( chunk_duration_secs: float = 300, # 5 minutes default overlap_secs: float = DEFAULT_OVERLAP_SECS, pause_search_secs: float = DEFAULT_PAUSE_SEARCH_SECS, + retry_gap_secs: float = DEFAULT_RETRY_GAP_SECS, ): """ Transcribe long audio by splitting into chunks and merging results. """ import nemo.collections.asr as nemo_asr - # Determine model path - model_filename = "parakeet-tdt-0.6b-v3.nemo" - model_path = None - # Locate project root: derived from VIRTUAL_ENV, which is set by `uv run` to path/.venv virtual_env = os.environ.get("VIRTUAL_ENV") if not virtual_env: @@ -199,10 +323,11 @@ def transcribe_buffered( sys.exit(1) project_root = os.path.dirname(virtual_env) - model_path = os.path.join(project_root, model_filename) + model_path = resolve_model_path(project_root) + model_name = os.path.splitext(os.path.basename(model_path))[0] if not os.path.exists(model_path): - print(f"Error during transcription: Can't find {model_filename} in project root: {project_root}") + print(f"Error during transcription: Can't find model file: {model_path}") sys.exit(1) print(f"Loading NVIDIA Parakeet model from: {model_path}") @@ -226,58 +351,25 @@ def transcribe_buffered( print(f"Splitting audio into chunks of at most {chunk_duration_secs}s " f"(overlap {overlap_secs}s, pause search {pause_search_secs}s)...") - chunks, sr, cut_times = split_audio_file( + chunks, sr, cut_times, audio = split_audio_file( audio_path, chunk_duration_secs, overlap_secs, pause_search_secs ) print(f"Created {len(chunks)} chunks") slice_results = [] chunk_texts = [] + retried = 0 + workdir = tempfile.mkdtemp(prefix="parakeet-") # per run: concurrent jobs share /tmp for i, chunk_info in enumerate(chunks): print(f"Transcribing chunk {i+1}/{len(chunks)} (duration: {chunk_info['duration']:.1f}s)...") - - # Save chunk to temporary file - chunk_path = f"/tmp/chunk_{i}.wav" - sf.write(chunk_path, chunk_info['audio'], sr) - - try: - # Transcribe chunk - output = asr_model.transcribe( - [chunk_path], - batch_size=1, - timestamps=True, - ) - - result_data = output[0] - chunk_text = result_data.text - chunk_texts.append(chunk_text) - chunk_words = [] - chunk_segments = [] - - # Extract and adjust timestamps - if hasattr(result_data, 'timestamp') and result_data.timestamp: - # Adjust timestamps by chunk start time - for word in result_data.timestamp.get("word", []): - word_copy = dict(word) - word_copy['start'] += chunk_info['start_time'] - word_copy['end'] += chunk_info['start_time'] - chunk_words.append(word_copy) - - for segment in result_data.timestamp.get("segment", []): - seg_copy = dict(segment) - seg_copy['start'] += chunk_info['start_time'] - seg_copy['end'] += chunk_info['start_time'] - chunk_segments.append(seg_copy) - - slice_results.append((chunk_words, chunk_segments)) - - print(f"Chunk {i+1} complete: {len(chunk_text)} characters") - - finally: - # Clean up temp file - if os.path.exists(chunk_path): - os.remove(chunk_path) + chunk_text, chunk_words, chunk_segments = transcribe_chunk( + asr_model, chunk_info['audio'], sr, chunk_info['start_time'], + os.path.join(workdir, f"chunk_{i}.wav"), + ) + chunk_texts.append(chunk_text) + slice_results.append((chunk_words, chunk_segments)) + print(f"Chunk {i+1} complete: {len(chunk_text)} characters") chunk_spans = [(c['start_time'], c['start_time'] + c['duration']) for c in chunks] all_words, all_segments = stitch_slices(slice_results, cut_times, chunk_spans) @@ -287,7 +379,20 @@ def transcribe_buffered( print("Warning: a chunk has text but no word timestamps; joining chunk texts") final_text = " ".join(chunk_texts) else: + if retry_gap_secs > 0: + def transcribe_span(start_sample, end_sample): + return transcribe_chunk(asr_model, audio[start_sample:end_sample], sr, + start_sample / sr, os.path.join(workdir, "retry.wav")) + try: + all_words, all_segments, retried = retry_speech_gaps( + all_words, all_segments, audio, sr, retry_gap_secs, transcribe_span + ) + except Exception as err: # best effort: never lose a finished transcription + print(f"Warning: gap retry failed ({err}); keeping the first pass") + if retried: + print(f"Re-transcribed {retried} stretch(es) of speech that got no words") final_text = " ".join(w["word"] for w in all_words) + shutil.rmtree(workdir, ignore_errors=True) print(f"Transcription complete: {len(final_text)} characters total") output_data = { @@ -296,13 +401,14 @@ def transcribe_buffered( "word_timestamps": all_words, "segment_timestamps": all_segments, "audio_file": audio_path, - "model": "parakeet-tdt-0.6b-v3", + "model": model_name, "buffered": True, "chunk_duration_secs": chunk_duration_secs, "num_chunks": len(chunks), "overlap_secs": overlap_secs, "pause_search_secs": pause_search_secs, "cut_times": cut_times, + "retried_gaps": retried, } if output_file: @@ -328,6 +434,12 @@ def main(): help=f"Seconds shared by adjacent chunks, capped at a quarter of --chunk-len " f"(default: {DEFAULT_OVERLAP_SECS}; 0 disables)" ) + parser.add_argument( + "--retry-gaps", type=float, default=DEFAULT_RETRY_GAP_SECS, + help=f"Re-transcribe on its own any stretch of at least this many seconds where " + f"the audio holds speech but the model produced no words " + f"(default: {DEFAULT_RETRY_GAP_SECS}; 0 disables)" + ) parser.add_argument( "--pause-search", type=float, default=DEFAULT_PAUSE_SEARCH_SECS, help=f"Seconds before each chunk limit searched for the quietest point to cut at, " @@ -346,6 +458,7 @@ def main(): chunk_duration_secs=args.chunk_len, overlap_secs=args.overlap, pause_search_secs=args.pause_search, + retry_gap_secs=args.retry_gaps, ) diff --git a/internal/transcription/adapters/py/nvidia/tests/test_parakeet_slicing.py b/internal/transcription/adapters/py/nvidia/tests/test_parakeet_slicing.py index 6a35947..c3d7537 100644 --- a/internal/transcription/adapters/py/nvidia/tests/test_parakeet_slicing.py +++ b/internal/transcription/adapters/py/nvidia/tests/test_parakeet_slicing.py @@ -10,7 +10,10 @@ import numpy as np import pytest sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) -from parakeet_transcribe_buffered import plan_slices, stitch_slices # noqa: E402 +from parakeet_transcribe_buffered import ( # noqa: E402 + find_speech_gaps, plan_slices, resolve_model_path, retry_speech_gaps, + segments_from_words, stitch_slices, +) SR = 16000 @@ -327,3 +330,157 @@ def test_punctuation_alone_is_never_an_anchor(): right = [word("-", 9.52, 9.55, 8), word("no", 10.4, 10.6, 8)] words, _ = stitch_slices([(left, []), (right, [])], [10.0], OVERLAPPING) assert words[0] is left[0] and texts(words) == ["-", "no"] + + +# -- model and attention selection ---------------------------------------------- + + +def test_the_bundled_v3_model_is_the_default(monkeypatch): + monkeypatch.delenv("PARAKEET_MODEL_PATH", raising=False) + assert resolve_model_path("/env") == "/env/parakeet-tdt-0.6b-v3.nemo" + + +def test_parakeet_model_path_overrides_the_model(monkeypatch): + monkeypatch.setenv("PARAKEET_MODEL_PATH", "/models/parakeet-tdt-0.6b-v2.nemo") + assert resolve_model_path("/env") == "/models/parakeet-tdt-0.6b-v2.nemo" + monkeypatch.setenv("PARAKEET_MODEL_PATH", "other.nemo") + assert resolve_model_path("/env") == "/env/other.nemo" + + +# -- retrying stretches where the model went quiet -------------------------------- + + +def talk(seconds, level=0.1, seed=3): + return (level * np.random.default_rng(seed).standard_normal(int(seconds * SR))).astype(np.float32) + + +def w_at(text, start, end): + return {"word": text, "start_offset": 0, "end_offset": 0, "start": start, "end": end} + + +def test_a_long_wordless_stretch_over_speech_is_a_gap(): + audio = talk(40) + words = [w_at("a", 1.0, 1.5), w_at("b", 9.5, 10.0), w_at("c", 20.0, 20.5), w_at("d", 38.0, 38.5)] + assert find_speech_gaps(words, audio, SR, min_gap=3.0) == [(1.5, 9.5), (10.0, 20.0), (20.5, 38.0)] + + +def test_a_wordless_stretch_over_silence_or_a_short_pause_is_not(): + audio = talk(40) + audio[int(10 * SR):int(20 * SR)] = 0.0 + words = [w_at("a", 1.0, 1.5), w_at("b", 9.5, 10.0), w_at("c", 20.0, 20.5), w_at("d", 22.0, 22.5), + w_at("e", 24.0, 24.5), w_at("f", 39.0, 39.5)] + # 10-20 s is silent; 20.5-22 and 22.5-24 are too short; 1.5-9.5 and 24.5-39 are talk + assert find_speech_gaps(words, audio, SR, min_gap=3.0) == [(1.5, 9.5), (24.5, 39.0)] + + +def test_speech_before_the_first_word_counts(): + audio = talk(20) + assert find_speech_gaps([w_at("late", 15.0, 15.5)], audio, SR, min_gap=3.0) == [(0.0, 15.0), (15.5, 20.0)] + + +def test_no_words_at_all_over_speech_is_one_gap(): + assert find_speech_gaps([], talk(12), SR, min_gap=3.0) == [(0.0, 12.0)] + + +def test_retried_words_are_spliced_into_the_gap_and_segments_rebuilt(): + audio = talk(31) + words = [w_at("Hello", 1.0, 1.5), w_at("there.", 2.0, 2.5), w_at("Bye.", 30.0, 30.5)] + segments = [segment(words[:2]), segment(words[2:])] + calls = [] + + def transcribe(s0, s1): # stands in for Parakeet: finds speech the main pass missed + calls.append((s0 / SR, s1 / SR)) + found = [w_at("We", 5.0, 5.2), w_at("missed", 6.0, 6.3), w_at("this.", 7.0, 7.4), + w_at("Again", 12.0, 12.4), w_at("outside", 40.5, 41.0)] + return "", [w for w in found if s0 / SR <= w["start"] < s1 / SR], [] + new_words, new_segments, n = retry_speech_gaps(words, segments, audio, SR, 3.0, transcribe) + assert [w["word"] for w in new_words] == ["Hello", "there.", "We", "missed", "this.", "Again", "Bye."] + assert n == 1 and calls == [(2.0, 30.5)] # one gap, 2.5-30.0 s, padded by 0.5 s + assert [g["segment"] for g in new_segments] == ["Hello there.", "We missed this.", "Again Bye."] + assert " ".join(g["segment"] for g in new_segments) == " ".join(w["word"] for w in new_words) + + +def test_without_gaps_nothing_is_retried_and_output_is_untouched(): + audio = talk(10) + words = [w_at("One", 0.5, 1.0), w_at("two", 2.5, 3.0), w_at("three.", 5.0, 5.5), w_at("four", 8.0, 8.5)] + segments = [segment(words[:3]), segment(words[3:])] + + def transcribe(s0, s1): + raise AssertionError("must not be called") + assert retry_speech_gaps(words, segments, audio, SR, 3.0, transcribe) == (words, segments, 0) + + +def test_long_gaps_are_retried_in_pieces(): + audio = talk(191) + words = [w_at("a", 1.0, 1.5), w_at("z", 190.0, 190.5)] + calls = [] + + def transcribe(s0, s1): + calls.append(round((s1 - s0) / SR, 1)) + return "", [], [] + retry_speech_gaps(words, [], audio, SR, 3.0, transcribe) + assert max(calls) <= 61.0 and len(calls) == 4 # 1.5-190 s in <= 60 s pieces (+0.5 s padding) + + +def test_segments_from_words_split_after_sentence_punctuation(): + ws = [w_at("Yes.", 0, 1), w_at("Is", 1, 2), w_at("it?", 2, 3), w_at("Go", 3, 4), w_at("now", 4, 5)] + assert [g["segment"] for g in segments_from_words(ws)] == ["Yes.", "Is it?", "Go now"] + assert segments_from_words([]) == [] + + +def test_a_retried_copy_of_the_word_after_the_gap_is_not_kept_twice(): + # The retry piece is padded, so it re-hears the next kept word and may + # timestamp it just inside the gap. + audio = talk(31) + words = [w_at("Hello.", 1.0, 1.5), w_at("Bye.", 30.0, 30.5)] + + def transcribe(s0, s1): + return "", [w_at("We", 5.0, 5.2), w_at("left.", 6.0, 6.3), w_at("bye.", 29.7, 30.3)], [] + new_words, _, _ = retry_speech_gaps(words, [], audio, SR, 3.0, transcribe) + assert [w["word"] for w in new_words] == ["Hello.", "We", "left.", "Bye."] + assert new_words[-1] is words[-1] + + +def test_a_word_heard_by_two_retry_pieces_is_kept_once(): + audio = talk(191) + words = [w_at("a", 1.0, 1.5), w_at("z", 190.0, 190.5)] + + def transcribe(s0, s1): # both pieces around 61.5 s hear "boundary" + t0, t1 = s0 / SR, s1 / SR + out = [w_at("boundary", 61.3, 61.8)] if t0 <= 61.3 < t1 else [] + out += [w_at("boundary", 61.55, 61.9)] if t0 <= 61.55 < t1 and t0 > 60 else [] + return "", out, [] + new_words, _, _ = retry_speech_gaps(words, [], audio, SR, 3.0, transcribe) + assert [w["word"] for w in new_words] == ["a", "boundary", "z"] + + +def test_a_failing_retry_keeps_the_first_pass(): + audio = talk(31) + words = [w_at("Hello.", 1.0, 1.5), w_at("Bye.", 30.0, 30.5)] + segments = [segment(words[:1]), segment(words[1:])] + + def transcribe(s0, s1): + raise RuntimeError("CUDA out of memory") + assert retry_speech_gaps(words, segments, audio, SR, 3.0, transcribe)[:2] == (words, segments) + + +def test_a_looping_retry_is_not_spliced_in(): + audio = talk(31) + words = [w_at("Hello.", 1.0, 1.5), w_at("Bye.", 30.0, 30.5)] + + def transcribe(s0, s1): # a hallucination loop over music-like audio + return "", [w_at("la", 3.0 + k, 3.2 + k) for k in range(20)], [] + new_words, _, _ = retry_speech_gaps(words, [], audio, SR, 3.0, transcribe) + assert [w["word"] for w in new_words] == ["Hello.", "Bye."] + + +def test_gap_levels_line_up_with_time_at_sample_rates_other_than_16k(): + # At 11025 Hz a 10 ms hop is 110 samples (9.977 ms); indexing frames by + # time/10ms would drift ~2 s by 900 s and read the silence before the gap. + sr = 11025 + audio = (0.1 * np.random.default_rng(5).standard_normal(1000 * sr)).astype(np.float32) + audio[int(896.9 * sr):int(900.0 * sr)] = 0.0 + words = [w_at("w", t, t + 0.5) for t in np.arange(0.0, 1000.0, 1.0) if not 900.0 <= t < 903.0] + words = [w for w in words if not (899.9 < w["start"] < 900.5)] + [w_at("w", 899.5, 900.0)] + words.sort(key=lambda w: w["start"]) + assert (900.0, 903.0) in find_speech_gaps(words, audio, sr, 3.0) -- 2.39.5