fix(hermes): reuse exact speculative speech snapshot

This commit is contained in:
jenkins 2026-08-24 01:28:48 -03:00
parent a23cf838b2
commit be5566cabd
2 changed files with 196 additions and 0 deletions

View File

@ -135,6 +135,59 @@ class InferenceGate:
INFERENCE_GATE = InferenceGate()
class StreamTelemetry:
"""Retain bounded evidence about the most recently committed voice turn."""
def __init__(self) -> None:
self._lock = threading.Lock()
self._speculative_cache_hits = 0
self._direct_final_decodes = 0
self._last_commit_cache_hit = False
self._last_commit_pending_wait_ms = 0.0
self._last_commit_total_ms = 0.0
self._last_commit_audio_ms = 0.0
def record_commit(
self,
*,
cache_hit: bool,
pending_wait_ms: float,
total_ms: float,
audio_bytes: int,
) -> None:
"""Record whether commit reused its exact speculative EOS snapshot."""
with self._lock:
if cache_hit:
self._speculative_cache_hits += 1
else:
self._direct_final_decodes += 1
self._last_commit_cache_hit = cache_hit
self._last_commit_pending_wait_ms = pending_wait_ms
self._last_commit_total_ms = total_ms
self._last_commit_audio_ms = (
audio_bytes
* 1000
/ (STREAM_SAMPLE_RATE * STREAM_SAMPLE_WIDTH)
)
def snapshot(self) -> dict:
"""Return the small aggregate used by health checks and release tests."""
with self._lock:
return {
"speculative_cache_hits": self._speculative_cache_hits,
"direct_final_decodes": self._direct_final_decodes,
"last_commit_cache_hit": self._last_commit_cache_hit,
"last_commit_pending_wait_ms": round(
self._last_commit_pending_wait_ms, 1
),
"last_commit_total_ms": round(self._last_commit_total_ms, 1),
"last_commit_audio_ms": round(self._last_commit_audio_ms, 1),
}
STREAM_TELEMETRY = StreamTelemetry()
STARTUP_WARMUP_MS: float | None = None
STARTUP_ROLLING_WARMUP_MS: float | None = None
@ -497,6 +550,8 @@ class StreamingTranscription:
self._pending_key = ""
self._cached_key = ""
self._cached_payload: dict | None = None
self._eos_pcm = b""
self._eos_key = ""
self._rolling_inflight = False
self._rolling_last_started_at = 0.0
self._rolling_speech_bytes = 0
@ -542,6 +597,8 @@ class StreamingTranscription:
frame is preserved until the next speculative EOS marker, even when a
quiet syllable falls below the server's conservative energy threshold.
"""
if self._at_eos and self._eos_pcm and self._eos_key:
return self._eos_pcm, self._eos_key
cutoff = len(self._pcm)
if self._heard_speech and self._last_speech_byte:
tail_bytes = STREAM_SAMPLE_RATE * STREAM_SAMPLE_WIDTH * VAD_TAIL_MS // 1000
@ -550,6 +607,11 @@ class StreamingTranscription:
key = f"{self._epoch}:{hashlib.sha256(pcm).hexdigest()}"
return pcm, key
def _clear_eos_snapshot_locked(self) -> None:
"""Invalidate frozen EOS audio when authoritative speech resumes."""
self._eos_pcm = b""
self._eos_key = ""
def _rolling_snapshot_locked(self) -> tuple[bytes, int]:
"""Return a bounded tail window and its absolute start in milliseconds."""
window_bytes = _pcm_bytes_for_ms(ROLLING_WINDOW_MS)
@ -693,6 +755,7 @@ class StreamingTranscription:
self._epoch += 1
self._cached_key = ""
self._cached_payload = None
self._clear_eos_snapshot_locked()
self._reset_rolling_history_locked(rearm=True)
self._heard_speech = True
self._at_eos = False
@ -728,6 +791,7 @@ class StreamingTranscription:
self._client_active = True
self._cached_key = ""
self._cached_payload = None
self._clear_eos_snapshot_locked()
self._reset_rolling_history_locked(rearm=True)
def speculate(self) -> None:
@ -738,6 +802,13 @@ class StreamingTranscription:
self._client_active = False
self._at_eos = True
pcm, key = self._snapshot_locked()
# The browser keeps forwarding PCM until recorder shutdown. Freeze
# the exact EOS snapshot so trailing silence cannot invalidate a
# completed Whisper-small speculative result. A resume message or
# newly detected speech clears this snapshot before more audio is
# accepted into the authoritative turn.
self._eos_pcm = pcm
self._eos_key = key
if not pcm or key in {self._pending_key, self._cached_key}:
return
self._pending_key = key
@ -794,6 +865,9 @@ class StreamingTranscription:
def commit(self) -> None:
"""Emit one final result from the complete stable utterance snapshot."""
commit_started = time.monotonic()
pending_wait_ms = 0.0
cache_hit = False
with self._lock:
if not self._started or self._closed or self._committing:
raise WebSocketError("turn cannot be committed", 1008)
@ -802,12 +876,15 @@ class StreamingTranscription:
self._committing = True
pcm, key = self._snapshot_locked()
deadline = time.monotonic() + STREAM_IDLE_SECONDS
pending_wait_started = time.monotonic()
while self._pending_key == key and not self._cached_payload:
remaining = deadline - time.monotonic()
if remaining <= 0:
break
self._lock.wait(timeout=remaining)
pending_wait_ms = (time.monotonic() - pending_wait_started) * 1000
payload = self._cached_payload if self._cached_key == key else None
cache_hit = payload is not None
try:
if payload is None:
@ -820,6 +897,12 @@ class StreamingTranscription:
model=payload["model"],
)
)
STREAM_TELEMETRY.record_commit(
cache_hit=cache_hit,
pending_wait_ms=pending_wait_ms,
total_ms=(time.monotonic() - commit_started) * 1000,
audio_bytes=len(pcm),
)
except Exception as exc:
print(f"[stt] streaming transcription failed: {exc}", flush=True)
self.error("transcription failed", fatal=True)
@ -834,6 +917,7 @@ class StreamingTranscription:
self._closed = True
self._cached_payload = None
self._cached_key = ""
self._clear_eos_snapshot_locked()
self._lock.notify_all()
@ -872,6 +956,7 @@ class SpeechHandler(BaseHTTPRequestHandler):
"rolling_model": ROLLING_MODEL_NAME,
"device": "cuda" if torch.cuda.is_available() else "cpu",
"inference": INFERENCE_GATE.snapshot(),
"stream_commits": STREAM_TELEMETRY.snapshot(),
"startup": {
"warmed": (
STARTUP_WARMUP_MS is not None

View File

@ -7,6 +7,7 @@ import struct
import sys
import threading
import time
import wave
from pathlib import Path
from types import SimpleNamespace
@ -80,6 +81,35 @@ class _Model:
}
class _BlockingModel(_Model):
"""Hold one authoritative decode to exercise commit/speculation races."""
def __init__(self, text: str = "hello") -> None:
super().__init__(text)
self.entered = threading.Event()
self.release = threading.Event()
def transcribe(self, path: str, **options: object) -> dict:
self.entered.set()
assert self.release.wait(2.0)
return super().transcribe(path, **options)
class _SequencedModel(_Model):
"""Return evolving text and retain the audio size decoded each time."""
def __init__(self, texts: list[str]) -> None:
super().__init__(texts[0])
self.texts = texts
self.frame_counts: list[int] = []
def transcribe(self, path: str, **options: object) -> dict:
with wave.open(path, "rb") as wav_file:
self.frame_counts.append(wav_file.getnframes())
self.text = self.texts[min(len(self.calls), len(self.texts) - 1)]
return super().transcribe(path, **options)
def _start(session) -> None:
session.start(
{
@ -96,6 +126,10 @@ def _speech(samples: int = 1_600) -> bytes:
return struct.pack(f"<{samples}h", *([12_000] * samples))
def _silence(samples: int = 12_000) -> bytes:
return bytes(samples * 2)
def _wait_idle(session) -> None:
deadline = time.monotonic() + 1.0
with session._lock:
@ -129,6 +163,83 @@ def test_rolling_uses_tiny_once_while_final_remains_small(monkeypatch):
assert len(final_model.calls) == 1
def test_trailing_silence_after_speculation_keeps_exact_cache_key(monkeypatch):
"""Recorder shutdown silence must not force a second Whisper-small decode."""
module = _load_server(monkeypatch)
connection = _Connection()
model = _Model("hello")
session = module.StreamingTranscription(connection, model)
_start(session)
session.append(_speech(samples=20_800))
session.speculate()
assert connection.wait_for("partial")["transcript"] == "hello"
session.append(_silence(samples=8_000))
session.commit()
assert connection.wait_for("final")["transcript"] == "hello"
assert len(model.calls) == 1
telemetry = module.STREAM_TELEMETRY.snapshot()
assert telemetry["speculative_cache_hits"] == 1
assert telemetry["direct_final_decodes"] == 0
assert telemetry["last_commit_cache_hit"] is True
assert telemetry["last_commit_audio_ms"] == 1_300.0
def test_commit_waits_for_one_inflight_exact_speculation_without_redecoding(
monkeypatch,
):
"""Exercise the browser timing race while Whisper-small is still running."""
module = _load_server(monkeypatch)
connection = _Connection()
model = _BlockingModel()
session = module.StreamingTranscription(connection, model)
_start(session)
session.append(_speech(samples=20_800))
session.speculate()
assert model.entered.wait(1.0)
session.append(_silence(samples=8_000))
committed = threading.Event()
worker = threading.Thread(target=lambda: (session.commit(), committed.set()))
worker.start()
time.sleep(0.02)
assert not committed.is_set()
model.release.set()
worker.join(timeout=1.0)
assert committed.is_set()
assert connection.wait_for("final")["transcript"] == "hello"
assert len(model.calls) == 1
telemetry = module.STREAM_TELEMETRY.snapshot()
assert telemetry["last_commit_cache_hit"] is True
assert telemetry["last_commit_pending_wait_ms"] >= 10
def test_resumed_speech_invalidates_frozen_eos_and_decodes_full_turn(monkeypatch):
"""Freezing EOS cannot discard an intentional barge-in continuation."""
module = _load_server(monkeypatch)
connection = _Connection()
model = _SequencedModel(["first thought", "complete thought"])
session = module.StreamingTranscription(connection, model)
_start(session)
session.append(_speech(samples=8_000))
session.speculate()
connection.wait_for("partial")
session.append(_silence(samples=4_000))
session.resume()
session.append(_speech(samples=4_000))
session.commit()
assert connection.wait_for("final")["transcript"] == "complete thought"
assert model.frame_counts == [8_000, 16_000]
telemetry = module.STREAM_TELEMETRY.snapshot()
assert telemetry["speculative_cache_hits"] == 0
assert telemetry["direct_final_decodes"] == 1
assert telemetry["last_commit_cache_hit"] is False
def test_eos_does_not_rearm_but_resumed_speech_does(monkeypatch):
module = _load_server(monkeypatch)
monkeypatch.setattr(module, "ROLLING_MIN_AUDIO_MS", 100)