Last active
June 4, 2026 02:20
-
-
Save akngs/1b1e1452a0526d66ceaa1c586546cc54 to your computer and use it in GitHub Desktop.
Omnispeak
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| local obj = {} | |
| obj.name = "OmniSpeak" | |
| obj.projectDir = "/Users/ak/prjs/prv/omnispeak" | |
| obj.uvPath = "/opt/homebrew/bin/uv" | |
| obj.zshPath = "/bin/zsh" | |
| obj.endpoint = "http://127.0.0.1:8765" | |
| obj.recording = false | |
| obj.serverTask = nil | |
| obj.serverStartedAt = nil | |
| obj.statusIcon = nil | |
| obj.waitingForServer = false | |
| obj.recoveringServer = false | |
| obj.recordingEmoji = "🐶" | |
| obj.faceCanvas = nil | |
| function obj:init() | |
| self.statusIcon = hs.menubar.new() | |
| self:setStatus("") | |
| return self | |
| end | |
| function obj:start() | |
| self:ensureServer() | |
| hs.timer.doAfter(2, function() | |
| self:ensureServer() | |
| end) | |
| return self | |
| end | |
| function obj:setStatus(text) | |
| if self.statusIcon then | |
| self.statusIcon:setTitle(text) | |
| end | |
| end | |
| function obj:bindHotkeys(mapping) | |
| local spec = { | |
| toggle = hs.fnutils.partial(self.toggle, self), | |
| } | |
| hs.spoons.bindHotkeysToSpec(spec, mapping) | |
| return self | |
| end | |
| function obj:decodeJson(body) | |
| local ok, decoded = pcall(hs.json.decode, body or "{}") | |
| if ok then | |
| return decoded | |
| end | |
| return nil | |
| end | |
| function obj:ensureServer() | |
| hs.http.asyncGet(self.endpoint .. "/health", nil, function(code, body) | |
| if code == 200 then | |
| local res = self:decodeJson(body) | |
| if res and res.ready == true then | |
| self:setStatus("") | |
| else | |
| self:setStatus("STT loading") | |
| self:waitForServer(0) | |
| end | |
| return | |
| end | |
| if self.serverTask and self.serverTask:isRunning() then | |
| local startedAt = self.serverStartedAt or hs.timer.secondsSinceEpoch() | |
| if hs.timer.secondsSinceEpoch() - startedAt < 10 then | |
| self:setStatus("STT loading") | |
| self:waitForServer(0) | |
| return | |
| end | |
| end | |
| self:recoverServer() | |
| end) | |
| end | |
| function obj:startServer() | |
| if self.serverTask and self.serverTask:isRunning() then | |
| return | |
| end | |
| self:setStatus("STT loading") | |
| self.serverStartedAt = hs.timer.secondsSinceEpoch() | |
| self.serverTask = hs.task.new(self.uvPath, function(exitCode, _stdout, stderr) | |
| self.serverTask = nil | |
| self.serverStartedAt = nil | |
| if exitCode ~= 0 then | |
| self:setStatus("") | |
| if not self.recoveringServer then | |
| hs.alert("OmniSpeak server exited: " .. tostring(stderr)) | |
| end | |
| end | |
| end, { | |
| "run", | |
| "--project", | |
| self.projectDir, | |
| "omnispeak", | |
| }) | |
| self.serverTask:start() | |
| self:waitForServer(0) | |
| end | |
| function obj:recoverServer() | |
| if self.recoveringServer then | |
| return | |
| end | |
| self.recoveringServer = true | |
| self:setStatus("STT starting") | |
| local script = string.format([==[ | |
| project=%q | |
| port_pids=$(lsof -tiTCP:8765 -sTCP:LISTEN 2>/dev/null || true) | |
| stale_pids="" | |
| blockers="" | |
| for pid in $port_pids; do | |
| [ "$pid" = "$$" ] && continue | |
| command=$(ps -p "$pid" -o command= 2>/dev/null || true) | |
| parent=$(ps -p "$pid" -o ppid= 2>/dev/null | tr -d " " || true) | |
| parent_command="" | |
| if [ -n "$parent" ]; then | |
| parent_command=$(ps -p "$parent" -o command= 2>/dev/null || true) | |
| fi | |
| if [[ "$command" == *"uv run --project ${project} omnispeak"* ]] || | |
| [[ "$command" == *"${project}/.venv/bin/omnispeak"* ]] || | |
| [[ "$parent_command" == *"uv run --project ${project} omnispeak"* ]]; then | |
| stale_pids="$stale_pids $pid" | |
| if [[ "$parent_command" == *"uv run --project ${project} omnispeak"* ]]; then | |
| stale_pids="$stale_pids $parent" | |
| fi | |
| else | |
| blockers="${blockers}\n${pid} ${command}" | |
| fi | |
| done | |
| if [ -n "$blockers" ]; then | |
| print -r -- "Port 8765 is already used by:${blockers}" >&2 | |
| exit 2 | |
| fi | |
| if [ -n "$stale_pids" ]; then | |
| /bin/kill -TERM $stale_pids 2>/dev/null || true | |
| sleep 1 | |
| for pid in $stale_pids; do | |
| if kill -0 "$pid" 2>/dev/null; then | |
| /bin/kill -KILL "$pid" 2>/dev/null || true | |
| fi | |
| done | |
| fi | |
| ]==], self.projectDir) | |
| hs.task.new(self.zshPath, function(exitCode, _stdout, stderr) | |
| self.recoveringServer = false | |
| if exitCode ~= 0 then | |
| self:setStatus("") | |
| hs.alert("OmniSpeak cleanup failed: " .. tostring(stderr)) | |
| return | |
| end | |
| self:startServer() | |
| end, { "-lc", script }):start() | |
| end | |
| function obj:waitForServer(attempt) | |
| if attempt == 0 then | |
| if self.waitingForServer then | |
| return | |
| end | |
| self.waitingForServer = true | |
| end | |
| if attempt > 120 then | |
| self.waitingForServer = false | |
| self:setStatus("") | |
| hs.alert("OmniSpeak server did not become ready") | |
| return | |
| end | |
| hs.timer.doAfter(1, function() | |
| hs.http.asyncGet(self.endpoint .. "/health", nil, function(code, body) | |
| local res = self:decodeJson(body) | |
| if code == 200 and res and res.ready == true then | |
| self.waitingForServer = false | |
| self:setStatus("") | |
| else | |
| local taskRunning = self.serverTask and self.serverTask:isRunning() | |
| if code ~= 200 and attempt >= 10 and not taskRunning then | |
| self.waitingForServer = false | |
| self:recoverServer() | |
| else | |
| self:waitForServer(attempt + 1) | |
| end | |
| end | |
| end) | |
| end) | |
| end | |
| function obj:toggle() | |
| self:ensureServer() | |
| if self.recording then | |
| self:stopRecording() | |
| else | |
| self:startRecording() | |
| end | |
| end | |
| function obj:startRecording() | |
| hs.http.asyncPost(self.endpoint .. "/start", "", { ["Content-Type"] = "application/json" }, function(code, body) | |
| if code ~= 200 then | |
| self:setStatus("") | |
| local res = self:decodeJson(body) | |
| local message = res and res.error or "OmniSpeak is not ready" | |
| hs.alert(tostring(message)) | |
| return | |
| end | |
| local res = self:decodeJson(body) | |
| if res and res.recording then | |
| self.recording = true | |
| self:setStatus("STT rec") | |
| self:showFace() | |
| end | |
| end) | |
| end | |
| function obj:stopRecording() | |
| local focusedWindow = hs.window.frontmostWindow() | |
| local focusedApp = focusedWindow and focusedWindow:application() or nil | |
| self:hideFace() | |
| self:setStatus("STT transcribing") | |
| hs.http.asyncPost(self.endpoint .. "/stop", "", { ["Content-Type"] = "application/json" }, function(code, body) | |
| self.recording = false | |
| self:setStatus("") | |
| if code ~= 200 then | |
| local res = self:decodeJson(body) | |
| local message = res and res.error or tostring(body) | |
| hs.alert("OmniSpeak failed: " .. tostring(message)) | |
| return | |
| end | |
| local res = self:decodeJson(body) | |
| local text = res and res.text or "" | |
| if text == "" then | |
| hs.alert("No speech detected") | |
| return | |
| end | |
| self:pasteText(text, focusedApp) | |
| end) | |
| end | |
| function obj:pasteText(text, focusedApp) | |
| local previousClipboard = hs.pasteboard.getContents() | |
| hs.pasteboard.setContents(text) | |
| if focusedApp then | |
| focusedApp:activate() | |
| end | |
| hs.timer.doAfter(0.05, function() | |
| hs.eventtap.keyStroke({ "cmd" }, "v") | |
| hs.timer.doAfter(0.25, function() | |
| if previousClipboard then | |
| hs.pasteboard.setContents(previousClipboard) | |
| end | |
| end) | |
| end) | |
| end | |
| function obj:showFace() | |
| if not self.faceCanvas then | |
| local screen = hs.screen.mainScreen() | |
| local frame = screen:frame() | |
| local size = 280 | |
| self.faceCanvas = hs.canvas.new({ | |
| x = frame.x + (frame.w - size) / 2, | |
| y = frame.y + (frame.h - size) / 2, | |
| w = size, | |
| h = size, | |
| }) | |
| self.faceCanvas:level(hs.canvas.windowLevels.overlay) | |
| self.faceCanvas:behavior({ "canJoinAllSpaces", "stationary" }) | |
| self.faceCanvas[1] = { | |
| type = "text", | |
| text = self.recordingEmoji, | |
| textSize = 200, | |
| textAlignment = "center", | |
| frame = { x = "0%", y = "5%", w = "100%", h = "95%" }, | |
| } | |
| else | |
| self.faceCanvas[1].text = self.recordingEmoji | |
| end | |
| self.faceCanvas:show() | |
| end | |
| function obj:hideFace() | |
| if self.faceCanvas then | |
| self.faceCanvas:hide() | |
| end | |
| end | |
| return obj |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| from __future__ import annotations | |
| import json | |
| import os | |
| import signal | |
| import sys | |
| import tempfile | |
| import threading | |
| import time | |
| from concurrent.futures import Future, ThreadPoolExecutor | |
| from http import HTTPStatus | |
| from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer | |
| from typing import Any, Callable | |
| import numpy as np | |
| import sounddevice as sd | |
| from mlx_audio.stt.generate import generate_transcription | |
| from mlx_audio.stt.utils import load_model | |
| from omnispeak.voiceprint import VOICEPRINT_PATH, load_encoder | |
| HOST = "127.0.0.1" | |
| PORT = 8765 | |
| MODEL_ID = "mlx-community/Qwen3-ASR-1.7B-4bit" | |
| LANGUAGE = "Korean" | |
| MAX_TOKENS = 256 | |
| SAMPLE_RATE = 16_000 | |
| VOICEPRINT_THRESHOLD = float(os.environ.get("OMNISPEAK_VOICEPRINT_THRESHOLD", "0.30")) | |
| class ModelNotReadyError(RuntimeError): | |
| pass | |
| class Recorder: | |
| def __init__(self) -> None: | |
| self.model: Any | None = None | |
| self.model_status = "not_loaded" | |
| self.model_error: str | None = None | |
| self.embed_fn: Callable[[np.ndarray], np.ndarray] | None = None | |
| self.reference: np.ndarray | None = None | |
| self.verifier_status = "not_loaded" | |
| self.verifier_error: str | None = None | |
| self.stream: sd.InputStream | None = None | |
| self.started_at: float | None = None | |
| self.chunks: list[np.ndarray] = [] | |
| self.lock = threading.Lock() | |
| self.executor = ThreadPoolExecutor( | |
| max_workers=1, | |
| thread_name_prefix="omnispeak-model", | |
| ) | |
| def load(self) -> None: | |
| with self.lock: | |
| self.model_status = "loading" | |
| self.model_error = None | |
| try: | |
| model = load_model(MODEL_ID) | |
| except Exception as exc: | |
| with self.lock: | |
| self.model_status = "error" | |
| self.model_error = str(exc) | |
| raise | |
| with self.lock: | |
| self.model = model | |
| self.model_status = "ready" | |
| self._load_verifier() | |
| def _load_verifier(self) -> None: | |
| with self.lock: | |
| self.verifier_status = "loading" | |
| self.verifier_error = None | |
| try: | |
| embed_fn = load_encoder() | |
| if VOICEPRINT_PATH.exists(): | |
| ref = np.load(VOICEPRINT_PATH).astype(np.float32) | |
| ref = ref / np.linalg.norm(ref) | |
| with self.lock: | |
| self.embed_fn = embed_fn | |
| self.reference = ref | |
| self.verifier_status = "ready" | |
| print( | |
| f"speaker verification enabled (threshold={VOICEPRINT_THRESHOLD})", | |
| file=sys.stderr, | |
| flush=True, | |
| ) | |
| else: | |
| with self.lock: | |
| self.embed_fn = embed_fn | |
| self.verifier_status = "no_voiceprint" | |
| print( | |
| f"no voiceprint at {VOICEPRINT_PATH}; speaker verification disabled", | |
| file=sys.stderr, | |
| flush=True, | |
| ) | |
| except Exception as exc: | |
| with self.lock: | |
| self.verifier_status = "error" | |
| self.verifier_error = str(exc) | |
| print(f"verifier load failed: {exc}", file=sys.stderr, flush=True) | |
| def load_async(self) -> Future[None]: | |
| return self.executor.submit(self.load) | |
| def close(self) -> None: | |
| self.executor.shutdown(wait=False, cancel_futures=True) | |
| @property | |
| def is_recording(self) -> bool: | |
| return self.stream is not None | |
| def start(self) -> dict[str, Any]: | |
| with self.lock: | |
| if self.model is None: | |
| if self.model_status == "error": | |
| message = f"model failed to load: {self.model_error}" | |
| else: | |
| message = "model is still loading" | |
| raise ModelNotReadyError(message) | |
| if self.is_recording: | |
| return {"recording": True, "already_recording": True} | |
| self.chunks = [] | |
| self.started_at = time.monotonic() | |
| def callback(indata: np.ndarray, _frames: int, _time: Any, status: Any) -> None: | |
| if status: | |
| print(status, file=sys.stderr, flush=True) | |
| with self.lock: | |
| self.chunks.append(indata[:, 0].copy()) | |
| self.stream = sd.InputStream( | |
| channels=1, | |
| samplerate=SAMPLE_RATE, | |
| dtype="float32", | |
| callback=callback, | |
| ) | |
| self.stream.start() | |
| return {"recording": True} | |
| def stop(self) -> dict[str, Any]: | |
| with self.lock: | |
| stream = self.stream | |
| chunks = self.chunks | |
| started_at = self.started_at | |
| self.stream = None | |
| self.chunks = [] | |
| self.started_at = None | |
| if stream is None: | |
| return {"recording": False, "text": "", "error": "not_recording"} | |
| stream.stop() | |
| stream.close() | |
| if not chunks: | |
| return {"recording": False, "text": "", "duration_seconds": 0.0} | |
| audio = np.concatenate(chunks).astype(np.float32) | |
| duration = len(audio) / SAMPLE_RATE | |
| if duration < 0.25: | |
| return { | |
| "recording": False, | |
| "text": "", | |
| "duration_seconds": duration, | |
| "samples": int(len(audio)), | |
| } | |
| if self.model is None: | |
| raise ModelNotReadyError("model is still loading") | |
| similarity: float | None = None | |
| if self.embed_fn is not None and self.reference is not None: | |
| similarity = self.executor.submit(self.verify_speaker, audio).result() | |
| if similarity < VOICEPRINT_THRESHOLD: | |
| print( | |
| f"speaker mismatch: similarity={similarity:.3f} <" | |
| f" threshold={VOICEPRINT_THRESHOLD}", | |
| file=sys.stderr, | |
| flush=True, | |
| ) | |
| return { | |
| "recording": False, | |
| "text": "", | |
| "duration_seconds": duration, | |
| "samples": int(len(audio)), | |
| "speaker_match": False, | |
| "similarity": similarity, | |
| } | |
| result = self.executor.submit(self.transcribe, audio).result() | |
| return { | |
| "recording": False, | |
| "text": str(getattr(result, "text", result)).strip(), | |
| "duration_seconds": duration, | |
| "elapsed_seconds": time.monotonic() - started_at if started_at else duration, | |
| "samples": int(len(audio)), | |
| "speaker_match": True if similarity is not None else None, | |
| "similarity": similarity, | |
| } | |
| def verify_speaker(self, audio: np.ndarray) -> float: | |
| if self.embed_fn is None or self.reference is None: | |
| raise ModelNotReadyError("verifier is not ready") | |
| vec = self.embed_fn(audio) | |
| vec = vec / np.linalg.norm(vec) | |
| return float(np.dot(self.reference, vec)) | |
| def transcribe(self, audio: np.ndarray) -> Any: | |
| if self.model is None: | |
| raise ModelNotReadyError("model is still loading") | |
| with tempfile.TemporaryDirectory(prefix="omnispeak-") as tmpdir: | |
| return generate_transcription( | |
| model=self.model, | |
| audio=audio, | |
| output_path=f"{tmpdir}/transcript", | |
| format="txt", | |
| verbose=False, | |
| language=LANGUAGE, | |
| max_tokens=MAX_TOKENS, | |
| ) | |
| def status(self) -> dict[str, Any]: | |
| with self.lock: | |
| elapsed = time.monotonic() - self.started_at if self.started_at else 0.0 | |
| status = { | |
| "recording": self.is_recording, | |
| "ready": self.model is not None, | |
| "model": MODEL_ID, | |
| "model_status": self.model_status, | |
| "language": LANGUAGE, | |
| "elapsed_seconds": elapsed, | |
| "verifier_status": self.verifier_status, | |
| "voiceprint_threshold": VOICEPRINT_THRESHOLD, | |
| } | |
| if self.model_error: | |
| status["model_error"] = self.model_error | |
| if self.verifier_error: | |
| status["verifier_error"] = self.verifier_error | |
| return status | |
| class OmniSpeakServer(ThreadingHTTPServer): | |
| allow_reuse_address = True | |
| daemon_threads = True | |
| recorder: Recorder | |
| class Handler(BaseHTTPRequestHandler): | |
| server: OmniSpeakServer | |
| def do_GET(self) -> None: | |
| if self.path in {"/health", "/status"}: | |
| self.write_json(HTTPStatus.OK, self.server.recorder.status()) | |
| return | |
| self.write_json(HTTPStatus.NOT_FOUND, {"error": "not_found"}) | |
| def do_POST(self) -> None: | |
| try: | |
| if self.path == "/start": | |
| self.write_json(HTTPStatus.OK, self.server.recorder.start()) | |
| return | |
| if self.path == "/stop": | |
| self.write_json(HTTPStatus.OK, self.server.recorder.stop()) | |
| return | |
| self.write_json(HTTPStatus.NOT_FOUND, {"error": "not_found"}) | |
| except ModelNotReadyError as exc: | |
| self.write_json(HTTPStatus.SERVICE_UNAVAILABLE, {"error": str(exc)}) | |
| except Exception as exc: | |
| self.write_json(HTTPStatus.INTERNAL_SERVER_ERROR, {"error": str(exc)}) | |
| def log_message(self, fmt: str, *args: Any) -> None: | |
| print(f"{self.address_string()} - {fmt % args}", file=sys.stderr, flush=True) | |
| def write_json(self, status: HTTPStatus, body: dict[str, Any]) -> None: | |
| payload = json.dumps(body).encode("utf-8") | |
| self.send_response(status) | |
| self.send_header("Content-Type", "application/json; charset=utf-8") | |
| self.send_header("Content-Length", str(len(payload))) | |
| self.end_headers() | |
| self.wfile.write(payload) | |
| def create_server( | |
| recorder: Recorder, | |
| address: tuple[str, int] = (HOST, PORT), | |
| ) -> OmniSpeakServer: | |
| server = OmniSpeakServer(address, Handler) | |
| server.recorder = recorder | |
| return server | |
| def start_model_loading(recorder: Recorder) -> Future[None]: | |
| print(f"loading model: {MODEL_ID}", file=sys.stderr, flush=True) | |
| future = recorder.load_async() | |
| def log_result(done: Future[None]) -> None: | |
| try: | |
| done.result() | |
| except Exception as exc: | |
| print(f"model load failed: {exc}", file=sys.stderr, flush=True) | |
| else: | |
| print("model loaded", file=sys.stderr, flush=True) | |
| future.add_done_callback(log_result) | |
| return future | |
| def main() -> int: | |
| recorder = Recorder() | |
| server = create_server(recorder) | |
| start_model_loading(recorder) | |
| def request_shutdown(_signum: int, _frame: Any) -> None: | |
| raise KeyboardInterrupt | |
| signal.signal(signal.SIGINT, request_shutdown) | |
| signal.signal(signal.SIGTERM, request_shutdown) | |
| print(f"listening on http://{HOST}:{PORT}", file=sys.stderr, flush=True) | |
| try: | |
| server.serve_forever() | |
| except KeyboardInterrupt: | |
| pass | |
| finally: | |
| server.server_close() | |
| recorder.close() | |
| return 0 | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| from __future__ import annotations | |
| import argparse | |
| import sys | |
| import time | |
| from pathlib import Path | |
| import numpy as np | |
| import sounddevice as sd | |
| SAMPLE_RATE = 16_000 | |
| DATA_DIR = Path.home() / ".omnispeak" | |
| VOICEPRINT_PATH = DATA_DIR / "voiceprint.npy" | |
| ENROLL_DIR = DATA_DIR / "enrollment" | |
| MODEL_DIR = DATA_DIR / "models" / "spkrec-ecapa-voxceleb" | |
| MODEL_SOURCE = "speechbrain/spkrec-ecapa-voxceleb" | |
| EMBED_DIM = 192 | |
| SAMPLES_PER_ENROLLMENT = 5 | |
| DURATION_SECONDS = 5.0 | |
| def load_encoder(): | |
| import torch | |
| from speechbrain.inference.classifiers import EncoderClassifier | |
| MODEL_DIR.mkdir(parents=True, exist_ok=True) | |
| classifier = EncoderClassifier.from_hparams( | |
| source=MODEL_SOURCE, | |
| savedir=str(MODEL_DIR), | |
| run_opts={"device": "cpu"}, | |
| ) | |
| def embed(audio: np.ndarray) -> np.ndarray: | |
| signal = torch.from_numpy(audio.astype(np.float32)).unsqueeze(0) | |
| with torch.inference_mode(): | |
| out = classifier.encode_batch(signal) | |
| vec = out.squeeze().cpu().numpy() | |
| vec = vec / np.linalg.norm(vec) | |
| return vec | |
| return embed | |
| def record(seconds: float) -> np.ndarray: | |
| print(f" recording {seconds:.1f}s... ", end="", flush=True) | |
| audio = sd.rec( | |
| int(seconds * SAMPLE_RATE), | |
| samplerate=SAMPLE_RATE, | |
| channels=1, | |
| dtype="float32", | |
| ) | |
| sd.wait() | |
| print("done") | |
| return audio[:, 0] | |
| def enroll( | |
| out_path: Path = VOICEPRINT_PATH, | |
| enroll_dir: Path = ENROLL_DIR, | |
| n: int = SAMPLES_PER_ENROLLMENT, | |
| ) -> None: | |
| embed = load_encoder() | |
| enroll_dir.mkdir(parents=True, exist_ok=True) | |
| for old in enroll_dir.glob("sample_*.npy"): | |
| old.unlink() | |
| print(f"\nEnrolling — {n} samples of {DURATION_SECONDS:.0f}s each.") | |
| print("Speak naturally; vary tone, speed, and phrasing between samples.\n") | |
| embeds = [] | |
| for i in range(n): | |
| input(f"[{i + 1}/{n}] Press Enter to start recording...") | |
| audio = record(DURATION_SECONDS) | |
| np.save(enroll_dir / f"sample_{i:02d}.npy", audio) | |
| vec = embed(audio) | |
| embeds.append(vec) | |
| print(f" embedding ok (dim={vec.shape[0]})") | |
| mean_embed = np.mean(np.stack(embeds), axis=0) | |
| mean_embed /= np.linalg.norm(mean_embed) | |
| out_path.parent.mkdir(parents=True, exist_ok=True) | |
| np.save(out_path, mean_embed) | |
| print(f"\nSaved voiceprint to {out_path}") | |
| print(f"Saved {n} raw recordings to {enroll_dir}") | |
| def extend( | |
| n: int, | |
| out_path: Path = VOICEPRINT_PATH, | |
| enroll_dir: Path = ENROLL_DIR, | |
| ) -> None: | |
| enroll_dir.mkdir(parents=True, exist_ok=True) | |
| existing = sorted(enroll_dir.glob("sample_*.npy")) | |
| if not existing: | |
| print( | |
| f"no saved recordings in {enroll_dir} — run enroll first", | |
| file=sys.stderr, | |
| ) | |
| sys.exit(1) | |
| start = int(existing[-1].stem.split("_")[1]) + 1 | |
| print(f"\nExtending — {n} more samples of {DURATION_SECONDS:.0f}s each.") | |
| print(f"Currently have {len(existing)} samples; will save sample_{start:02d}..sample_{start + n - 1:02d}.\n") | |
| for i in range(n): | |
| input(f"[{i + 1}/{n}] Press Enter to start recording...") | |
| audio = record(DURATION_SECONDS) | |
| np.save(enroll_dir / f"sample_{start + i:02d}.npy", audio) | |
| print(f" saved sample_{start + i:02d}.npy") | |
| reembed(out_path=out_path, enroll_dir=enroll_dir) | |
| def reembed( | |
| out_path: Path = VOICEPRINT_PATH, | |
| enroll_dir: Path = ENROLL_DIR, | |
| ) -> None: | |
| samples = sorted(enroll_dir.glob("sample_*.npy")) | |
| if not samples: | |
| print( | |
| f"no saved recordings in {enroll_dir} — run enroll first", | |
| file=sys.stderr, | |
| ) | |
| sys.exit(1) | |
| embed = load_encoder() | |
| print(f"\nRe-embedding {len(samples)} saved recordings...") | |
| embeds = [] | |
| for path in samples: | |
| audio = np.load(path) | |
| vec = embed(audio) | |
| embeds.append(vec) | |
| print(f" {path.name} ok") | |
| mean_embed = np.mean(np.stack(embeds), axis=0) | |
| mean_embed /= np.linalg.norm(mean_embed) | |
| out_path.parent.mkdir(parents=True, exist_ok=True) | |
| np.save(out_path, mean_embed) | |
| print(f"\nSaved voiceprint to {out_path}") | |
| def verify(in_path: Path = VOICEPRINT_PATH) -> None: | |
| if not in_path.exists(): | |
| print( | |
| f"no voiceprint at {in_path} — run `omnispeak-voiceprint enroll` first", | |
| file=sys.stderr, | |
| ) | |
| sys.exit(1) | |
| reference = np.load(in_path) | |
| if reference.shape[0] != EMBED_DIM: | |
| print( | |
| f"voiceprint dim={reference.shape[0]} does not match ECAPA dim={EMBED_DIM};" | |
| " re-enroll with `omnispeak-voiceprint enroll`", | |
| file=sys.stderr, | |
| ) | |
| sys.exit(1) | |
| reference = reference / np.linalg.norm(reference) | |
| embed = load_encoder() | |
| print(f"\nLoaded reference from {in_path}") | |
| print("Press Ctrl+C to quit.\n") | |
| while True: | |
| try: | |
| input("Press Enter to record a verification sample...") | |
| except (EOFError, KeyboardInterrupt): | |
| print() | |
| return | |
| audio = record(DURATION_SECONDS) | |
| t0 = time.monotonic() | |
| vec = embed(audio) | |
| elapsed_ms = (time.monotonic() - t0) * 1000 | |
| similarity = float(np.dot(reference, vec)) | |
| print(f" similarity: {similarity:+.3f}") | |
| print(f" embedding latency: {elapsed_ms:.1f} ms\n") | |
| def main() -> int: | |
| parser = argparse.ArgumentParser( | |
| description="Voiceprint enrollment and verification (ECAPA-TDNN).", | |
| ) | |
| sub = parser.add_subparsers(dest="cmd", required=True) | |
| sub.add_parser("enroll", help="Record samples and save a voiceprint.") | |
| sub.add_parser("verify", help="Record samples and compare against the voiceprint.") | |
| sub.add_parser( | |
| "reembed", | |
| help="Recompute the voiceprint from saved recordings (no re-recording).", | |
| ) | |
| extend_parser = sub.add_parser( | |
| "extend", | |
| help="Record N more samples on top of existing ones and recompute the voiceprint.", | |
| ) | |
| extend_parser.add_argument("--n", type=int, default=2, help="Number of samples to add (default: 2).") | |
| args = parser.parse_args() | |
| if args.cmd == "enroll": | |
| enroll() | |
| elif args.cmd == "verify": | |
| verify() | |
| elif args.cmd == "reembed": | |
| reembed() | |
| elif args.cmd == "extend": | |
| extend(n=args.n) | |
| return 0 | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) |
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment