Hands-free voice mode had no language signal at all, so every spoken reply was synthesized with the English voice no matter what the user actually said. The multilingual Piper work (PR #26) added server-side routing for a "language" field but nothing ever sent one. Carry the language the private Jetson Whisper service already detects through to the TTS request for the reply that speech produced, and only for that reply. hermes-stt returns {text, model, language}, accepted only as a bare ISO-639 token; hermes_stt_client.py writes a <stem>.language sidecar next to the .txt transcript Hermes reads, leaving the local-command contract intact; the patched local-command envelope and /api/transcribe re-validate it and surface it; atlas-voice.js binds it to the voice-mode generation token and chat session, consumes it exactly once, and clears it on cancellation, restart, session change, empty transcript or transcription error; /api/tts honours it only from the fixed en/ru/es allow-list and otherwise sends English. A client "voice" field is never read at any hop, and typed messages, the manual read-aloud button, and any reply not produced by a spoken turn carry no trusted signal and stay on the English voice. The two WebUI-side and one agent-side edits are fail-closed replace_exact patches; both patch roots are now env-overridable so the contract can be verified offline without a GPU or an image build.
206 lines
7.4 KiB
Python
206 lines
7.4 KiB
Python
#!/usr/bin/env python3
|
|
"""Small OpenAI-compatible Whisper service for the dedicated Jetson."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import cgi
|
|
import json
|
|
import os
|
|
import re
|
|
import tempfile
|
|
import threading
|
|
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
|
from pathlib import Path
|
|
|
|
import torch
|
|
import whisper
|
|
|
|
|
|
HOST = os.getenv("HERMES_STT_HOST", "0.0.0.0")
|
|
PORT = int(os.getenv("HERMES_STT_PORT", "9000"))
|
|
MODEL_NAME = os.getenv("HERMES_STT_MODEL", "small")
|
|
CACHE_DIR = Path(os.getenv("HERMES_STT_CACHE", "/cache/whisper"))
|
|
MAX_AUDIO_BYTES = 30 * 1024 * 1024
|
|
MODEL_LOCK = threading.Lock()
|
|
|
|
|
|
def _repetitive_token(token: str) -> bool:
|
|
"""Identify long periodic Whisper hallucinations caused by steady noise."""
|
|
letters = "".join(re.findall(r"[a-z]+", token.lower()))
|
|
if len(letters) < 10:
|
|
return False
|
|
for period in range(1, 5):
|
|
pattern = letters[:period]
|
|
matches = sum(
|
|
character == pattern[index % period]
|
|
for index, character in enumerate(letters)
|
|
)
|
|
if matches / len(letters) >= 0.86:
|
|
return True
|
|
return max(letters.count(character) for character in set(letters)) / len(letters) >= 0.78
|
|
|
|
|
|
def _clean_transcript(result: dict) -> str:
|
|
"""Drop noise-only segments and repetitive tokens while retaining speech."""
|
|
segments = result.get("segments")
|
|
if not isinstance(segments, list):
|
|
segments = [{"text": result.get("text") or ""}]
|
|
kept: list[str] = []
|
|
for segment in segments:
|
|
if not isinstance(segment, dict):
|
|
continue
|
|
text = str(segment.get("text") or "").strip()
|
|
if not text:
|
|
continue
|
|
no_speech = float(segment.get("no_speech_prob") or 0.0)
|
|
average_logprob = float(segment.get("avg_logprob") or 0.0)
|
|
if no_speech >= 0.55 and average_logprob <= -0.55:
|
|
continue
|
|
words = [word for word in text.split() if not _repetitive_token(word)]
|
|
if words:
|
|
kept.append(" ".join(words))
|
|
return " ".join(kept).strip()
|
|
|
|
|
|
def _detected_language(result: object) -> str:
|
|
"""Return the bare ISO-639 code Whisper decoded with, or nothing at all.
|
|
|
|
``whisper.transcribe`` reports the language it auto-detected (or the one it
|
|
was told to use) as a plain lowercase token such as ``en``/``ru``/``yue``.
|
|
Anything that is not that exact shape is dropped rather than guessed at, so
|
|
a surprising model result can never become a downstream voice selector.
|
|
"""
|
|
if not isinstance(result, dict):
|
|
return ""
|
|
value = result.get("language")
|
|
if not isinstance(value, str):
|
|
return ""
|
|
code = value.strip().lower()
|
|
if not 2 <= len(code) <= 3 or not code.isascii() or not code.isalpha():
|
|
return ""
|
|
return code
|
|
|
|
|
|
def _transcription_payload(result: dict) -> dict:
|
|
"""Build the transcription contract: text plus the model's own language."""
|
|
return {
|
|
"text": _clean_transcript(result),
|
|
"model": MODEL_NAME,
|
|
"language": _detected_language(result),
|
|
}
|
|
|
|
|
|
def _json(handler: BaseHTTPRequestHandler, status: int, payload: dict) -> None:
|
|
body = json.dumps(payload).encode("utf-8")
|
|
handler.send_response(status)
|
|
handler.send_header("Content-Type", "application/json")
|
|
handler.send_header("Content-Length", str(len(body)))
|
|
handler.send_header("Cache-Control", "no-store")
|
|
handler.end_headers()
|
|
handler.wfile.write(body)
|
|
|
|
|
|
class SpeechHandler(BaseHTTPRequestHandler):
|
|
"""Serve health and transcription without exposing a general runtime."""
|
|
|
|
server_version = "AtlasWhisper/1"
|
|
|
|
def log_message(self, message: str, *args: object) -> None:
|
|
print(f"[stt] {self.address_string()} {message % args}", flush=True)
|
|
|
|
def do_GET(self) -> None:
|
|
if self.path != "/health":
|
|
_json(self, 404, {"error": "not found"})
|
|
return
|
|
_json(
|
|
self,
|
|
200,
|
|
{
|
|
"ok": True,
|
|
"model": MODEL_NAME,
|
|
"device": "cuda" if torch.cuda.is_available() else "cpu",
|
|
},
|
|
)
|
|
|
|
def do_POST(self) -> None:
|
|
if self.path != "/v1/audio/transcriptions":
|
|
_json(self, 404, {"error": "not found"})
|
|
return
|
|
content_length = int(self.headers.get("Content-Length", "0") or "0")
|
|
if content_length <= 0 or content_length > MAX_AUDIO_BYTES:
|
|
_json(self, 413, {"error": "audio payload is missing or too large"})
|
|
return
|
|
|
|
content_type = self.headers.get("Content-Type", "")
|
|
if not content_type.lower().startswith("multipart/form-data"):
|
|
_json(self, 400, {"error": "multipart/form-data is required"})
|
|
return
|
|
|
|
form = cgi.FieldStorage(
|
|
fp=self.rfile,
|
|
headers=self.headers,
|
|
environ={
|
|
"REQUEST_METHOD": "POST",
|
|
"CONTENT_TYPE": content_type,
|
|
"CONTENT_LENGTH": str(content_length),
|
|
},
|
|
)
|
|
audio = form["file"] if "file" in form else None
|
|
if audio is None or not getattr(audio, "file", None):
|
|
_json(self, 400, {"error": "file is required"})
|
|
return
|
|
|
|
suffix = Path(getattr(audio, "filename", "audio.wav") or "audio.wav").suffix
|
|
suffix = suffix if suffix in {".wav", ".webm", ".ogg", ".mp3", ".m4a"} else ".wav"
|
|
language = str(form.getfirst("language", "auto") or "auto").strip().lower()
|
|
temp_path = ""
|
|
try:
|
|
with tempfile.NamedTemporaryFile(prefix="atlas-stt-", suffix=suffix, delete=False) as temp:
|
|
temp_path = temp.name
|
|
while True:
|
|
chunk = audio.file.read(1024 * 1024)
|
|
if not chunk:
|
|
break
|
|
temp.write(chunk)
|
|
|
|
with MODEL_LOCK:
|
|
result = self.server.model.transcribe( # type: ignore[attr-defined]
|
|
temp_path,
|
|
language=None if language in {"", "auto"} else language,
|
|
task="transcribe",
|
|
fp16=torch.cuda.is_available(),
|
|
condition_on_previous_text=False,
|
|
temperature=0,
|
|
compression_ratio_threshold=2.0,
|
|
logprob_threshold=-0.8,
|
|
no_speech_threshold=0.5,
|
|
verbose=False,
|
|
)
|
|
_json(self, 200, _transcription_payload(result))
|
|
except Exception as exc:
|
|
print(f"[stt] transcription failed: {exc}", flush=True)
|
|
_json(self, 500, {"error": "transcription failed"})
|
|
finally:
|
|
if temp_path:
|
|
try:
|
|
os.unlink(temp_path)
|
|
except OSError:
|
|
pass
|
|
|
|
|
|
def main() -> None:
|
|
"""Warm Whisper once, then serve concurrent clients through one GPU lock."""
|
|
CACHE_DIR.mkdir(parents=True, exist_ok=True)
|
|
if not torch.cuda.is_available():
|
|
raise RuntimeError("CUDA is required for the Atlas Whisper service")
|
|
print(f"[stt] loading Whisper {MODEL_NAME} into CUDA", flush=True)
|
|
model = whisper.load_model(MODEL_NAME, device="cuda", download_root=str(CACHE_DIR))
|
|
server = ThreadingHTTPServer((HOST, PORT), SpeechHandler)
|
|
server.model = model # type: ignore[attr-defined]
|
|
print(f"[stt] ready on {HOST}:{PORT}", flush=True)
|
|
server.serve_forever(poll_interval=0.25)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|