-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathserver.py
More file actions
3036 lines (2738 loc) · 132 KB
/
Copy pathserver.py
File metadata and controls
3036 lines (2738 loc) · 132 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
835
836
837
838
839
840
841
842
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
862
863
864
865
866
867
868
869
870
871
872
873
874
875
876
877
878
879
880
881
882
883
884
885
886
887
888
889
890
891
892
893
894
895
896
897
898
899
900
901
902
903
904
905
906
907
908
909
910
911
912
913
914
915
916
917
918
919
920
921
922
923
924
925
926
927
928
929
930
931
932
933
934
935
936
937
938
939
940
941
942
943
944
945
946
947
948
949
950
951
952
953
954
955
956
957
958
959
960
961
962
963
964
965
966
967
968
969
970
971
972
973
974
975
976
977
978
979
980
981
982
983
984
985
986
987
988
989
990
991
992
993
994
995
996
997
998
999
1000
#!/usr/bin/env python3
"""
ASR server with both HTTP batch and WebSocket streaming transcription.
HTTP (batch):
POST /v1/audio/transcriptions — OpenAI-compatible, multipart file upload
WebSocket (streaming):
WS /ws/transcribe — send binary PCM16 16kHz mono frames, receive JSON:
{"partial": "text so far..."} — interim result while speaking
{"final": "complete sentence"} — after silence detected
{"done": true} — server closed stream
Health:
GET /health
"""
import gc
import os
import sys
import io
import json
import time
import uuid
import struct
import asyncio
import tempfile
import logging
import wave
import threading
from datetime import datetime
from pathlib import Path
from typing import Optional
import mlx.core as mx
import numpy as np
import uvicorn
from fastapi import FastAPI, File, Form, Body, UploadFile, HTTPException, WebSocket, WebSocketDisconnect
from fastapi.responses import JSONResponse, PlainTextResponse
# Shared poly* modules live alongside this MLX server in the repo root.
_REPO_ROOT = str(Path(__file__).resolve().parent)
if _REPO_ROOT not in sys.path:
sys.path.append(_REPO_ROOT)
from livestack_node import (ModelManager as AsrModelManager, ManagedUnit, ResidencyPolicy, # noqa: E402
counting, free_mlx, trim_ram)
import polyasr_align # noqa: E402
import polyasr_diarize # noqa: E402
from asr_evidence import AsrEvidence, inventory as evidence_inventory # noqa: E402
os.environ["PATH"] = "/opt/homebrew/bin:" + os.environ.get("PATH", "")
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s [%(levelname)s] %(message)s",
)
log = logging.getLogger("polyasr")
def _env(name: str, default: str) -> str:
"""Read POLYASR_<name>, falling back to the legacy ASR_<name> so existing
launchd units keep working until they're updated."""
return os.environ.get(f"POLYASR_{name}") or os.environ.get(f"ASR_{name}", default)
def _envflag(name: str, default: str) -> bool:
return _env(name, default).lower() not in {"0", "false", "no"}
MODEL_NAME = _env("MODEL", "Qwen/Qwen3-ASR-0.6B")
MODEL_REVISION = "5eb144179a02acc5e5ba31e748d22b0cf3e303b0"
_asr_evidence = AsrEvidence()
# MLX forced-alignment models (loaded lazily by /v1/align). The aligner needs an
# ASR model + a separate forced-aligner model, both via mlx_audio.stt.
ALIGN_ASR_MODEL = _env("ALIGN_ASR_MODEL", "mlx-community/Qwen3-ASR-0.6B-4bit")
ALIGN_ALIGNER_MODEL = _env("ALIGN_ALIGNER_MODEL", "mlx-community/Qwen3-ForcedAligner-0.6B-4bit")
# Speaker-diarization unit (pyannote, loaded lazily by /v1/diarize). pyannote on
# Apple Silicon MPS is flaky and Metal memory isn't returned to the OS on evict,
# so default to CPU — idle-evict then actually reclaims via gc + malloc_trim.
DIARIZE_DEVICE = _env("DIARIZE_DEVICE", "cpu")
# Idle-evict the resident model after this many seconds (0 = never evict) so a
# co-resident polytts/renderer can reclaim the GPU/Metal memory.
IDLE_EVICT_SECONDS = int(_env("IDLE_EVICT_SECONDS", "180"))
# Co-residency: when on (default), loading a unit does NOT evict the other, so
# benchday's 'asr' stays warm even when unchain loads 'align'. Set to 0 for the
# historical one-model-resident eviction behaviour.
COLOAD = _envflag("COLOAD", "1")
# Livestack node identity: this host's id in the residence broker/planner.
HOST_ID = _env("HOST_ID", "xc-mac-studio")
_vad_model = None
_voice_encoder = None
_transcribe_lock = threading.Lock()
# Single dedicated GPU worker thread. MLX's Metal stream is THREAD-LOCAL: the
# stream is created on whichever thread first evaluates a graph, and any later
# op on a different thread raises "There is no Stream(gpu, N) in current
# thread" (and can crash the process with a Metal "Invalid Resource" abort).
# asyncio's default executor (run_in_executor(None, ...)) is a multi-thread
# pool, so transcribe/load/warmup calls were landing on different threads and
# corrupting the stream. Every MLX call — model load, warmup, partials, finals,
# native streaming, unload, alignment — must run on this ONE thread (matches
# the polycore.ModelManager design note: "all load/unload calls should run on a
# single GPU executor thread, exactly like polytts").
from concurrent.futures import ThreadPoolExecutor as _ThreadPoolExecutor
_gpu_executor = _ThreadPoolExecutor(max_workers=1, thread_name_prefix="gpu")
# How busy this node is, reported to the broker and to client pickers via
# `attach(in_flight=)` -> `/livestack/capability` `load.in_flight`. It has to be
# counted here rather than derived from livestack leases: transcription reaches
# the model through `manager.ensure`, which takes a `__usage__:` lease that
# marks recency rather than work and lives for the whole idle-evict TTL. Counting
# those would report the node busy for the whole window after one request; not
# counting them (which is correct) leaves real concurrent work invisible.
#
# The measure is GPU-QUEUE DEPTH, not request count: `_gpu_executor` has exactly
# one worker, so everything awaiting it is either running or queued behind the
# one that is. That is the number a consumer's picker actually wants — it is what
# `MeshRoutePicker` normalizes as `in_flight` — and unlike a request count it
# also sees keep-warm and eviction work, which occupies the card just as much.
_busy = counting()
async def _run_on_gpu(fn, *args, **kwargs):
"""Await a blocking MLX call on the single GPU worker thread."""
loop = asyncio.get_event_loop()
if kwargs:
from functools import partial as _partial
fn = _partial(fn, **kwargs)
async with _busy:
return await loop.run_in_executor(_gpu_executor, fn, *args)
# Monotonic timestamp of the most recent transcription, used by the keep-warm
# loop to skip warmups when real traffic is already keeping the model hot.
_last_transcribe_monotonic = 0.0
# Last transcribe that produced NON-EMPTY text. Real successful traffic is
# freshness evidence in its own right: the probe only runs in idle gaps, so
# continuous (healthy) dictation would otherwise starve it and read as stale.
_last_good_transcribe_monotonic = 0.0
# MLX allocates Metal buffers aggressively and never returns them to the OS.
# Cap decoder length (ASR utterances rarely need >256 tokens) and clear the
# cache after every transcription so memory stays bounded.
ASR_MAX_NEW_TOKENS = int(_env("MAX_NEW_TOKENS", "256"))
# A short or near-silent clip must never run the decoder out to the full token
# budget. Low-signal audio makes ASR models loop/hallucinate until
# max_new_tokens, which on a contended MLX GPU can take 20-30s — long enough to
# blow past the client's final-wait timeout. Worse, because every inference
# (partials, finals, batch) serializes on _transcribe_lock, one runaway stalls
# all of them behind it. Cap tokens by audio length; 25 tok/s is 4-8x real
# speech, so legitimate transcripts are never truncated (long clips still hit
# the 256 ceiling unchanged — only short clips get a tighter, safer bound).
ASR_MIN_NEW_TOKENS = int(_env("MIN_NEW_TOKENS", "32"))
ASR_TOKENS_PER_AUDIO_SEC = int(_env("TOKENS_PER_AUDIO_SEC", "25"))
# Keep the MLX model/Metal kernels hot. After an idle gap the first inference
# takes ~9s (vs ~0.6s warm) — long enough that the first streaming partial
# misses the client's ~10s no-partial timeout, so live text never appears on
# the first dictation after a lull (only the batch fallback delivers, at stop).
# A tiny periodic warmup during idle gaps keeps it resident. 0 disables.
ASR_KEEPWARM_INTERVAL_SEC = float(_env("KEEPWARM_INTERVAL_SEC", "45"))
ASR_NATIVE_STREAMING = _envflag("NATIVE_STREAMING", "0")
ASR_FINAL_WAIT_PARTIAL = _envflag("FINAL_WAIT_PARTIAL", "0")
ASR_PARTIALS_ENABLED = _envflag("PARTIALS_ENABLED", "0")
ASR_STREAM_CHUNK_SEC = float(_env("STREAM_CHUNK_SEC", "2.0"))
ASR_STREAM_MAX_CONTEXT_SEC = float(_env("STREAM_MAX_CONTEXT_SEC", "30.0"))
ASR_STREAM_FINALIZATION_MODE = _env("STREAM_FINALIZATION_MODE", "latency")
# Streaming-quality self-heal. The persistent MLX Session's short-window decode
# path silently rots over long uptime (days): it returns EMPTY text for real
# speech on the streaming/partial path while the full-buffer (batch/final) path
# and /health stay green — so Harmony's residency+liveness health never trips,
# and dictation dies with the model still "healthy". The keepwarm loop already
# runs an idle-gated inference through the SAME _transcribe_buffer the partials
# use; we upgrade it to transcribe a canonical speech clip and assert non-empty
# output. Consecutive blanks => reload the asr unit through the (Harmony-wired)
# manager; if a reload doesn't clear it, exit for launchd to restart the process.
ASR_STREAMING_PROBE = _envflag("STREAMING_PROBE", "1")
# Consecutive blank probes before reloading the asr unit in place.
ASR_PROBE_FAIL_RELOAD = int(_env("PROBE_FAIL_RELOAD", "3"))
# Consecutive blank probes (i.e. a reload didn't help) before exiting the
# process so launchd/keepalive restarts it. 0 disables the exit escalation.
ASR_PROBE_FAIL_EXIT = int(_env("PROBE_FAIL_EXIT", "6"))
_PROBE_WAV_PATH = Path(__file__).parent / "assets" / "probe_speech_16k.wav"
# Live rot detection from real dictation traffic (complements the idle probe:
# catches the rot on the next real session instead of waiting ~2min of idle
# probing). The safe signature is the SAME divergence we diagnosed, not "a
# partial was empty" (which is legitimately empty for noise/onsets/low-confidence
# audio): a session that carried sustained speech signal AND whose FINAL
# succeeded (proving there was transcribable speech) yet emitted ZERO non-empty
# partials. `LIVE_RELOAD_WITNESSES` consecutive such sessions (any healthy
# partial resets the count) flags a reload, executed by the idle-gated keepwarm
# loop — never mid-stream.
ASR_LIVE_ROT_DETECT = _envflag("LIVE_ROT_DETECT", "1")
# Min seconds of signal-bearing audio for a blank-partials session to count as a
# rot witness (below this, blank partials are inconclusive, not evidence).
ASR_LIVE_MIN_SPEECH_SEC = float(_env("LIVE_MIN_SPEECH_SEC", "2.0"))
# Consecutive rot-witness sessions before requesting a reload (hysteresis).
ASR_LIVE_RELOAD_WITNESSES = int(_env("LIVE_RELOAD_WITNESSES", "2"))
# A probe result that is merely OLD must not read as healthy: on 2026-07-14 the
# probe stopped running entirely (evicted model → keepwarm `continue`d) and
# /health kept serving `ok: true` with a stale `last_text` straight through a
# total outage. A successful probe must land within this window while the model
# is servable, or health is degraded. Default = 4 keepwarm intervals.
ASR_PROBE_STALE_AFTER_SEC = float(_env("PROBE_STALE_AFTER_SEC", "0")) or None
# How often to verify the model is loadable at all (cheap, no GPU, does not
# resurrect an evicted unit).
ASR_WEIGHTS_CHECK_INTERVAL_SEC = float(_env("WEIGHTS_CHECK_INTERVAL_SEC", "60"))
# -------------------------------------------------------------------------
# Session logging: audio + events are archived per-session for troubleshooting
# (VAD tuning, speaker-embedding diagnostics, ASR regression tests). Disabled
# by setting ASR_LOG_DIR="". Raw PCM is written incrementally (crash-safe)
# and converted to FLAC (lossless, ~60% of WAV size) at session close.
# -------------------------------------------------------------------------
_log_dir_env = _env("LOG_DIR", "logs")
if _log_dir_env:
_p = Path(_log_dir_env)
LOG_DIR = _p if _p.is_absolute() else Path(__file__).parent / _p
else:
LOG_DIR = None
# Silero VAD: chunk must be exactly 512 samples at 16kHz (~32ms).
VAD_CHUNK_SAMPLES = 512
VAD_THRESHOLD = 0.5
# Minimum audio length (sec) before computing a speaker embedding. Embeddings
# on very short clips are unstable.
MIN_EMBED_SEC = 1.0
MIN_EMBED_BYTES_CONST = int(MIN_EMBED_SEC * 16000 * 2)
# Cosine similarity threshold for "same speaker as reference". Was 0.70, which
# rejected the account owner's own voice: over 7 days (2026-09-16..23) every
# rejected score sat between 0.59 and 0.70 — 69 rejects, 259 s of speech — and
# each one froze coverage and sent the dictation to a full batch re-upload.
SPEAKER_SIM_THRESHOLD = float(_env("SPEAKER_SIM_THRESHOLD", "0.50"))
ASR_PROTOCOL_VERSION = 1
ASR_FRAME_MAGIC = b"BASR"
ASR_FRAME_HEADER_BYTES = 16
ASR_FRAME_TYPE_AUDIO = 1
ASR_RESUME_TTL_SEC = float(_env("RESUME_TTL_SEC", "300"))
# Control-message protocol. The binary audio FRAMING is unchanged (still
# ASR_PROTOCOL_VERSION), so a v2 server keeps accepting v1 frames byte for byte;
# v2 only adds `finalCapturedSeq` on stop and `recognizedThroughSeq` on final.
# Negotiated per session: a client that asks for v2 gets `control` echoed in the
# ack, and one that does not stays on v1 semantics.
ASR_CONTROL_PROTOCOL_VERSION = 2
# How long a completed stop result stays answerable so a reconnecting client
# gets the SAME transcript instead of provoking a second finalization of one
# utterance. Past this horizon the server says so explicitly (`stopExpired`)
# rather than silently decoding again under a stop id it no longer recognizes.
ASR_STOP_RESULT_TTL_SEC = float(_env("STOP_RESULT_TTL_SEC", "300"))
# How much uncommitted accepted audio must pile up before the server commits a
# stable recognized prefix. The commit boundary is always a VAD silence
# boundary (that is the only thing that appends to `gated_audio`), so segments
# never cut a word and no acoustic overlap has to be carried into the live tail
# — the bounded overlap the design allows for is zero by construction of the
# boundary, which is also why the merge needs no text-level dedup guess.
ASR_STABLE_COMMIT_MIN_SEC = float(_env("STABLE_COMMIT_MIN_SEC", "8"))
def join_transcript(prefix: str, suffix: str) -> str:
"""Concatenate two independently recognized segments.
Chinese/Japanese/Korean do not use interword spaces, so a separator between
two CJK characters is a visible defect; everywhere else one space is right.
Decided on the two characters at the seam — never by guessing whether the
segments overlap in text, which is the mistake that makes concatenated
recognition drop or duplicate words.
(`_join_text` further down in the MLX server is pre-existing dead code with
the same intent and no callers; it predates this change and is left for its
owner to remove.)
"""
prefix = (prefix or "").strip()
suffix = (suffix or "").strip()
if not prefix:
return suffix
if not suffix:
return prefix
sep = "" if _is_cjk(prefix[-1]) and _is_cjk(suffix[0]) else " "
return f"{prefix}{sep}{suffix}"
def _is_cjk(ch: str) -> bool:
cp = ord(ch)
return (
0x4E00 <= cp <= 0x9FFF
or 0x3040 <= cp <= 0x30FF
or 0xAC00 <= cp <= 0xD7AF
)
class AsrProtocolSession:
"""In-memory journal for one ASR streaming utterance."""
def __init__(self, session_id: str):
now = time.monotonic()
self.session_id = session_id
self.created = now
self.updated = now
self.chunks = {}
self.highest_contiguous_seq = -1
self.gated_audio = bytearray()
self.pending_audio = bytearray()
self.partial_audio = bytearray()
self.raw_partial_audio = bytearray()
self.staging = bytearray()
self.gate_state = {'silence_run': 0, 'gate_was_open': False}
self.reference_embedding = None
self.last_partial_text = ""
self.raw_signal_bytes = 0
self.raw_signal_bytes_at_last_partial = 0
self.native_stream_state = None
# The connection's CoverageLedger. Kept here so a resumed connection
# continues it: a fresh ledger on resume could only claim the bytes the
# NEW socket carried, so every final after a mid-dictation reconnect
# read as uncovered and sent the client to a full re-upload.
self.ledger = None
self.final_text = None
self.final_stop_id = None
# Coverage proof for the cached final: the highest contiguous client
# chunk that transcript actually includes. None means "this text was not
# produced by decoding the accepted audio" (a promoted partial), which a
# v2 client must not treat as complete.
self.final_recognized_through_seq = None
self.final_at = None
# Immutable recognized prefix and the byte offset into `gated_audio` it
# covers. Everything before `stable_bytes` has been decoded once and is
# never decoded again — this is what stops a stop from costing a full
# re-transcription of the whole utterance.
self.stable_text = ""
self.stable_bytes = 0
def accept(self, seq: int, payload: bytes) -> bool:
self.updated = time.monotonic()
if seq in self.chunks:
return False
self.chunks[seq] = payload
while self.highest_contiguous_seq + 1 in self.chunks:
self.highest_contiguous_seq += 1
return True
def raw_audio(self) -> bytes:
return b"".join(self.chunks[seq] for seq in sorted(self.chunks))
def sync_from_connection(
self,
*,
gated_audio: bytearray,
pending_audio: bytearray,
partial_audio: bytearray,
raw_partial_audio: bytearray,
staging: bytearray,
gate_state: dict,
reference_embedding,
last_partial_text: str,
raw_signal_bytes: int,
raw_signal_bytes_at_last_partial: int,
native_stream_state,
stable_text: str = "",
stable_bytes: int = 0,
) -> None:
self.updated = time.monotonic()
self.gated_audio = bytearray(gated_audio)
self.pending_audio = bytearray(pending_audio)
self.partial_audio = bytearray(partial_audio)
self.raw_partial_audio = bytearray(raw_partial_audio)
self.staging = bytearray(staging)
self.gate_state = dict(gate_state)
self.reference_embedding = reference_embedding
self.last_partial_text = last_partial_text
self.raw_signal_bytes = raw_signal_bytes
self.raw_signal_bytes_at_last_partial = raw_signal_bytes_at_last_partial
self.native_stream_state = native_stream_state
self.stable_text = stable_text
self.stable_bytes = stable_bytes
def stop_result_is_fresh(self, stop_id: str) -> bool:
"""Can this stop id still be answered from the cache?"""
if self.final_text is None or self.final_stop_id != stop_id:
return False
if self.final_at is None:
return False
return (time.monotonic() - self.final_at) <= ASR_STOP_RESULT_TTL_SEC
_protocol_sessions = {}
def _prune_protocol_sessions():
now = time.monotonic()
stale = [
sid for sid, sess in _protocol_sessions.items()
if now - sess.updated > ASR_RESUME_TTL_SEC
]
for sid in stale:
_protocol_sessions.pop(sid, None)
def _clear_mlx_cache() -> None:
"""Force Python GC then drop unused MLX Metal buffers back to the OS."""
gc.collect()
mx.metal.clear_cache()
def _new_protocol_session(session_id: str) -> AsrProtocolSession:
_prune_protocol_sessions()
sess = AsrProtocolSession(session_id)
_protocol_sessions[session_id] = sess
return sess
def _get_or_create_protocol_session(session_id: str) -> AsrProtocolSession:
_prune_protocol_sessions()
sess = _protocol_sessions.get(session_id)
if sess is None:
sess = AsrProtocolSession(session_id)
_protocol_sessions[session_id] = sess
return sess
def _decode_protocol_audio_frame(frame: bytes):
if len(frame) < ASR_FRAME_HEADER_BYTES:
return None
if frame[:4] != ASR_FRAME_MAGIC:
return None
version = frame[4]
frame_type = frame[5]
if version != ASR_PROTOCOL_VERSION or frame_type != ASR_FRAME_TYPE_AUDIO:
return None
seq = int.from_bytes(frame[8:16], "big", signed=False)
return seq, frame[ASR_FRAME_HEADER_BYTES:]
def compute_health(
*,
weights_present: bool,
probe_enabled: bool,
consecutive_failures: int,
fail_threshold: int,
last_ok_age_sec: Optional[float],
stale_after_sec: Optional[float],
cache_unreachable: bool = False,
) -> tuple[str, list[str]]:
"""Derive (status, reasons) — the honest health of the streaming path.
PURE: no I/O, no clock (the age is passed in), so the incident state is
directly testable. `/health.status` was previously the literal "ok", which
made the fleet's only paging signal a constant.
Degraded when:
- the weights cannot be loaded at all (a wiped HF cache; the model then
transcribes to EMPTY, which looks exactly like silence), or
- the streaming probe is failing (the silent partial-window rot), or
- no probe has SUCCEEDED within the staleness window. Absence of evidence
is not health: a probe that stopped running kept reporting green
(consecutive_failures == 0) through a total outage.
"""
reasons: list[str] = []
if cache_unreachable:
reasons.append("model_cache_unreachable")
if not weights_present:
reasons.append("weights_missing")
if probe_enabled:
if fail_threshold > 0 and consecutive_failures >= fail_threshold:
reasons.append("streaming_probe_failing")
if stale_after_sec and (last_ok_age_sec is None or last_ok_age_sec > stale_after_sec):
reasons.append("streaming_probe_stale")
return ("ok" if not reasons else "degraded", reasons)
class ModelUnavailable(RuntimeError):
"""The ASR model could not be loaded (weights missing, OOM, bad checkout).
Distinct from an inference error on a loaded model: an unloadable model must
surface as a hard failure, never as an empty transcript. See
docs/health-honesty.md.
"""
# Set when the model cache cannot be opened at all. See _probe_model_cache_access.
_cache_access_error: Optional[str] = None
CACHE_ACCESS_TIMEOUT_SEC = float(_env("CACHE_ACCESS_TIMEOUT_SEC", "15"))
def _probe_model_cache_access() -> Optional[str]:
"""Can this process OPEN files in the HF cache? Bounded; returns the error.
On macOS a launchd-started process that opens a file on an external volume
(xc-mac-studio: /opt/xc-data, where the home-offload symlinks point) blocks
in the kernel waiting for a consent approval nobody can give:
mac_vnode_check_open -> Sandbox approval_solicit
-> __WAITING_ON_APPROVAL_FROM_SANDBOXD__
It never returns and never errors, so the model load hung forever with the
service "running". Processes started over ssh inherit sshd's Full Disk
Access, which is why it never reproduced by hand. This turns that silent
wedge into a named, logged refusal. docs/mlx-host-restart-wedge.md.
"""
from huggingface_hub import constants
hub = constants.HF_HUB_CACHE
done = threading.Event()
failure: list[str] = []
def touch():
try:
for name in os.listdir(hub):
path = os.path.join(hub, name)
if os.path.isfile(path):
with open(path, "rb") as f:
f.read(1)
break
except OSError as e:
failure.append(repr(e))
finally:
done.set()
# Daemon thread: if the open is stuck in the kernel it stays stuck, but it
# no longer holds the model lock or the GPU thread.
threading.Thread(target=touch, name="cache-access-probe", daemon=True).start()
if not done.wait(CACHE_ACCESS_TIMEOUT_SEC):
return (f"opening files in {os.path.realpath(hub)} did not return in "
f"{CACHE_ACCESS_TIMEOUT_SEC:.0f}s. On macOS this is the sandbox "
"waiting for consent to an external volume: keep HF_HOME on the "
"internal disk for launchd services")
return failure[0] if failure else None
def _asr_weights_present() -> bool:
"""Can the ASR model be loaded from the local HF cache, right now?
Cheap: resolves the snapshot offline — no GPU, no inference, and it does NOT
resurrect an evicted unit (so it is safe to call on a timer while Harmony has
the model evicted). This is the check that catches a wiped/incomplete cache,
which is otherwise invisible until the next reload.
"""
if os.path.isdir(MODEL_NAME): # a local path, not an HF repo id
return True
try:
from huggingface_hub import snapshot_download
snapshot_download(
MODEL_NAME,
allow_patterns=["*.json", "*.safetensors", "*.txt", "*.model"],
local_files_only=True,
)
return True
except Exception:
return False
def _load_asr_model():
"""Loader for the streaming/batch ASR model (managed unit 'asr')."""
if _cache_access_error:
raise ModelUnavailable(f"model cache unreachable: {_cache_access_error}")
log.info("Loading model %s ...", MODEL_NAME)
import mlx_qwen3_asr
session = mlx_qwen3_asr.Session(MODEL_NAME)
log.info("ASR model loaded successfully.")
return session
def _load_align_models():
"""Loader for the forced-alignment unit 'align': an MLX ASR model plus a
separate MLX forced-aligner model (both via mlx_audio.stt.load_model).
Returns (asr_model, aligner_model)."""
log.info("Loading MLX aligner: %s + %s ...", ALIGN_ASR_MODEL, ALIGN_ALIGNER_MODEL)
from mlx_audio.stt.utils import load_model as load_stt_model
asr_model = load_stt_model(ALIGN_ASR_MODEL)
aligner_model = load_stt_model(ALIGN_ALIGNER_MODEL)
log.info("MLX aligner models loaded.")
return (asr_model, aligner_model)
def _load_diarize_model():
"""Loader for the speaker-diarization unit 'diarize': a pyannote pipeline
(shared loader so the MLX and CUDA servers stay DRY). Runs on DIARIZE_DEVICE
(CPU by default on Apple Silicon)."""
log.info("Loading diarization pipeline: %s (device=%s) ...",
polyasr_diarize.DEFAULT_MODEL, DIARIZE_DEVICE)
pipeline = polyasr_diarize.load_pipeline(DIARIZE_DEVICE)
log.info("Diarization pipeline loaded.")
return pipeline
# Model manager. With COLOAD on (default) 'asr' (streaming/batch) and 'align'
# (forced alignment) co-reside so benchday's asr is never evicted by an unchain
# align; with COLOAD off they never co-reside. The 'diarize' unit (pyannote) is
# normally NOT loaded — it loads lazily on the first /v1/diarize call and is
# idle-evicted like the others. Either way idle-evict reclaims memory for a
# co-resident polytts/renderer.
# footprint = Metal bytes (weights + peak activation), mirrors the CUDA node's
# estimates; refine with livestack_node.measure_footprint(). 'asr' is HARD_PIN so
# the broker never evicts the primary dictation model for a co-resident TTS.
_UNITS = {
"asr": ManagedUnit("asr", _load_asr_model, free_mlx, footprint=5_000_000_000,
residency_policy=ResidencyPolicy.HARD_PIN),
"align": ManagedUnit("align", _load_align_models, free_mlx, footprint=3_000_000_000),
"diarize": ManagedUnit("diarize", _load_diarize_model, free_mlx, footprint=2_500_000_000),
}
# The manager + residency policy (LocalCoordinator standalone, or livestack's
# LivestackCoordinator) is wired by attach() once `app` exists, below.
manager = None # polycore.ModelManager, set by attach() / fallback below
residence = None # LivestackCoordinator (None in standalone mode)
def get_session():
"""Return the resident ASR session, loading it (evicting any other unit) if
needed. Resets the idle timer. Call from within _transcribe_lock."""
return manager.ensure("asr")
def get_vad():
"""Lazy-load Silero VAD (ONNX)."""
global _vad_model
if _vad_model is None:
log.info("Loading Silero VAD ...")
from silero_vad import load_silero_vad
_vad_model = load_silero_vad(onnx=True)
log.info("Silero VAD loaded.")
return _vad_model
def get_encoder():
"""Lazy-load Resemblyzer voice encoder."""
global _voice_encoder
if _voice_encoder is None:
log.info("Loading Resemblyzer voice encoder ...")
from resemblyzer import VoiceEncoder
_voice_encoder = VoiceEncoder(device="cpu", verbose=False)
log.info("Voice encoder loaded.")
return _voice_encoder
def pcm_to_float32(pcm: bytes) -> np.ndarray:
"""Convert PCM16 bytes to normalized float32 mono samples."""
if len(pcm) == 0:
return np.zeros(0, dtype=np.float32)
return np.frombuffer(pcm, dtype=np.int16).astype(np.float32) / 32768.0
def pcm_has_signal(pcm: bytes) -> bool:
"""Cheap energy gate for raw partial scheduling.
Silero VAD can miss short or quiet tails, so this intentionally uses a
low bar. It only suppresses obvious silence from repeatedly triggering
expensive raw-buffer transcription.
"""
if len(pcm) < 2:
return False
samples = np.frombuffer(pcm, dtype=np.int16)
if samples.size == 0:
return False
abs_samples = np.abs(samples.astype(np.int32))
return float(abs_samples.mean()) >= 80 or int(abs_samples.max()) >= 900
def vad_speech_prob(pcm_bytes: bytes) -> float:
"""Max Silero-VAD speech probability across the chunks in *pcm_bytes*."""
samples = pcm_to_float32(pcm_bytes)
if len(samples) < VAD_CHUNK_SAMPLES:
return 0.0
import torch
model = get_vad()
max_prob = 0.0
for i in range(0, len(samples) - VAD_CHUNK_SAMPLES + 1, VAD_CHUNK_SAMPLES):
chunk = torch.from_numpy(samples[i:i + VAD_CHUNK_SAMPLES])
prob = model(chunk, 16000).item()
if prob > max_prob:
max_prob = prob
return max_prob
def compute_embedding(pcm_bytes: bytes) -> Optional[np.ndarray]:
"""Compute a speaker embedding for the given PCM16 audio.
Returns None if the clip is too short or preprocessing fails.
"""
if len(pcm_bytes) < MIN_EMBED_BYTES_CONST:
return None
try:
from resemblyzer import preprocess_wav
wav = pcm_to_float32(pcm_bytes)
wav = preprocess_wav(wav, source_sr=16000)
if len(wav) < 16000: # preprocess_wav trims silence; need >=1s
return None
enc = get_encoder()
return enc.embed_utterance(wav)
except Exception:
log.exception("Embedding failed")
return None
def cosine_sim(a: np.ndarray, b: np.ndarray) -> float:
denom = float(np.linalg.norm(a) * np.linalg.norm(b))
if denom == 0.0:
return 0.0
return float(np.dot(a, b) / denom)
class SessionLogger:
"""Per-connection logger. Writes raw PCM to *input.pcm* as audio arrives
(append-only, crash-safe) and events to *events.jsonl*. On close, PCM is
transcoded to lossless FLAC and the .pcm file is removed.
All methods swallow exceptions — logging failures must not break ASR.
"""
def __init__(self, kind: str = "ws"):
self.enabled = False
self.dir: Optional[Path] = None
self.pcm_file = None
self.events_file = None
self.start_monotonic = time.monotonic()
self.session_id = uuid.uuid4().hex[:8]
self.bytes_written = 0
if LOG_DIR is None:
return
try:
now = datetime.now()
day_dir = LOG_DIR / "sessions" / now.strftime("%Y-%m-%d")
day_dir.mkdir(parents=True, exist_ok=True)
self.dir = day_dir / f"{now.strftime('%H%M%S')}-{kind}-{self.session_id}"
self.dir.mkdir()
self.pcm_file = open(self.dir / "input.pcm", "wb")
self.events_file = open(self.dir / "events.jsonl", "w", encoding="utf-8")
self.enabled = True
self.event("start", {"session_id": self.session_id, "kind": kind,
"model": MODEL_NAME})
except Exception:
log.exception("SessionLogger init failed; logging disabled for this session")
self.enabled = False
def _ms(self) -> int:
return int((time.monotonic() - self.start_monotonic) * 1000)
def audio(self, data: bytes) -> None:
if not self.enabled:
return
try:
self.pcm_file.write(data)
self.bytes_written += len(data)
except Exception:
log.exception("SessionLogger.audio write failed")
def event(self, type_: str, data: Optional[dict] = None) -> None:
if not self.enabled:
return
try:
ev = {"t_ms": self._ms(), "type": type_}
if data:
ev.update(data)
self.events_file.write(json.dumps(ev, ensure_ascii=False) + "\n")
self.events_file.flush()
except Exception:
log.exception("SessionLogger.event write failed")
def close(self) -> None:
"""Close handles and transcode PCM → FLAC. Safe to call once."""
if not self.enabled:
return
# Logged BEFORE disabling: event() is a no-op once disabled, so the
# other order never wrote `close` or the session's duration.
self.event("close", {"audio_bytes": self.bytes_written,
"duration_ms": self._ms()})
self.enabled = False
try:
if self.pcm_file:
self.pcm_file.close()
if self.events_file:
self.events_file.close()
except Exception:
log.exception("SessionLogger close failed")
# Transcode PCM → FLAC (lossless, ~40% smaller than WAV).
pcm_path = self.dir / "input.pcm"
flac_path = self.dir / "input.flac"
try:
if pcm_path.exists() and pcm_path.stat().st_size > 0:
import soundfile as sf
pcm = np.fromfile(str(pcm_path), dtype=np.int16)
sf.write(str(flac_path), pcm, 16000,
format="FLAC", subtype="PCM_16")
pcm_path.unlink()
log.info("Session log: %s (%.2fs audio)", self.dir,
len(pcm) / 16000.0)
elif pcm_path.exists():
pcm_path.unlink() # empty file — drop it
except Exception:
log.exception("FLAC transcode failed; keeping .pcm")
def _log_http_request(audio_bytes: bytes, filename: str, text: str,
language: Optional[str]) -> None:
"""Archive an HTTP batch request + its transcription."""
if LOG_DIR is None:
return
try:
now = datetime.now()
day_dir = LOG_DIR / "http" / now.strftime("%Y-%m-%d")
day_dir.mkdir(parents=True, exist_ok=True)
rid = uuid.uuid4().hex[:8]
prefix = day_dir / f"{now.strftime('%H%M%S')}-{rid}"
suffix = Path(filename).suffix if filename else ".bin"
(prefix.with_suffix(suffix)).write_bytes(audio_bytes)
meta = {"timestamp": now.isoformat(timespec="seconds"),
"filename": filename,
"language": language,
"model": MODEL_NAME,
"text": text}
(prefix.with_suffix(".json")).write_text(
json.dumps(meta, ensure_ascii=False, indent=2), encoding="utf-8")
except Exception:
log.exception("HTTP request logging failed")
def pcm_to_wav_bytes(pcm_data: bytes, sample_rate: int = 16000) -> bytes:
"""Wrap raw PCM16 mono data in a WAV container."""
buf = io.BytesIO()
with wave.open(buf, "wb") as wf:
wf.setnchannels(1)
wf.setsampwidth(2)
wf.setframerate(sample_rate)
wf.writeframes(pcm_data)
return buf.getvalue()
def rms_energy(pcm_data: bytes) -> float:
"""Calculate RMS energy of PCM16 data."""
n_samples = len(pcm_data) // 2
if n_samples == 0:
return 0.0
samples = struct.unpack(f"<{n_samples}h", pcm_data[:n_samples * 2])
return (sum(s * s for s in samples) / n_samples) ** 0.5
def pcm16_to_float32(pcm_data: bytes) -> np.ndarray:
"""Convert raw PCM16 mono bytes into the native streaming API format."""
if not pcm_data:
return np.zeros(0, dtype=np.float32)
return np.frombuffer(pcm_data, dtype=np.int16).astype(np.float32) / 32768.0
def native_streaming_available() -> bool:
if not ASR_NATIVE_STREAMING:
return False
with _transcribe_lock:
session = get_session()
return all(
hasattr(session, name)
for name in ("init_streaming", "feed_audio", "finish_streaming")
)
def init_native_stream(context: str):
with _transcribe_lock:
session = get_session()
return session.init_streaming(
context=context or "",
chunk_size_sec=ASR_STREAM_CHUNK_SEC,
max_context_sec=ASR_STREAM_MAX_CONTEXT_SEC,
max_new_tokens=ASR_MAX_NEW_TOKENS,
finalization_mode=ASR_STREAM_FINALIZATION_MODE,
)
async def feed_native_stream(state, pcm_data: bytes) -> str:
if state is None or not pcm_data:
return ""
pcm = pcm16_to_float32(pcm_data)
if pcm.size == 0:
return getattr(state, "text", "") or ""
loop = asyncio.get_event_loop()
def run_feed():
with _transcribe_lock:
return get_session().feed_audio(pcm, state)
await _run_on_gpu(run_feed)
return (getattr(state, "text", "") or "").strip()
async def finish_native_stream(state) -> str:
if state is None:
return ""
loop = asyncio.get_event_loop()
def run_finish():
with _transcribe_lock:
return get_session().finish_streaming(state)
await _run_on_gpu(run_finish)
_clear_mlx_cache()
return (getattr(state, "text", "") or "").strip()
# ---------------------------------------------------------------------------
app = FastAPI(title="polyasr (MLX)", version="2.0.0")
def _gpu_call(fn):
"""Run a thunk on the single MLX GPU executor thread (facade warm/evict),
the same thread all transcribe/load/unload calls use — MLX's Metal stream
is thread-local, so residence mutations must run here too."""
return _gpu_executor.submit(fn).result()
# Become a livestack node: lease-driven residence exposed under /livestack, so a
# host-broker can arbitrate Metal VRAM across this ASR server and co-resident
# TTS/renderer. Without livestack_node, polycore's LocalCoordinator reproduces the
# standalone COLOAD + idle-evict behaviour unchanged.
try:
from livestack_node import attach
# `port` is what lets this node report for duty to the host broker on its
# own — no LIVESTACK_PEERS entry, no broker restart. It is the one fact the
# broker cannot infer (a POST shows it our source address, not what we
# listen on), and it is the same value we hand uvicorn at the bottom.
def _inventory() -> dict:
return evidence_inventory(backend="qwen-mlx", model=MODEL_NAME,
model_revision=MODEL_REVISION, streaming="synthetic",
evidence=_asr_evidence)
manager, residence = attach(app, host_id=HOST_ID, kind="polyasr", units=_UNITS,
idle_seconds=IDLE_EVICT_SECONDS, coload=COLOAD,
gpu_call=_gpu_call, port=int(_env("PORT", "8765")),
in_flight=_busy, inventory=_inventory)
log.info("livestack residence attached (host=%s, kind=polyasr)", HOST_ID)
except ImportError:
manager = AsrModelManager(_UNITS, IDLE_EVICT_SECONDS, coload=COLOAD)
log.info("livestack_node absent; standalone LocalCoordinator")
@app.on_event("startup")
async def startup_event():
global _cache_access_error
_cache_access_error = _probe_model_cache_access()
if _cache_access_error:
log.critical("MODEL CACHE UNREACHABLE — ASR will not load: %s",
_cache_access_error)
await _refresh_weights_state(force=True)
return
log.info("Model cache access: ok")
log.info("Pre-loading ASR model at startup (idle_evict=%ss)...", IDLE_EVICT_SECONDS)
# Load the model ON the GPU worker thread so the MLX Metal stream is created
# there, not on the event-loop thread. Every later eval also runs on this
# thread, so the stream is always found (see _gpu_executor).
def _load_asr():
with _transcribe_lock:
get_session()
await _run_on_gpu(_load_asr)
get_vad()
get_encoder()
log.info(
"ASR config: partials=%s partial_interval=%.2fs partial_min_delta=%.2fs native_streaming=%s final_wait_partial=%s",
ASR_PARTIALS_ENABLED,
PARTIAL_INTERVAL_SEC,
PARTIAL_MIN_DELTA_SEC,
ASR_NATIVE_STREAMING,
ASR_FINAL_WAIT_PARTIAL,
)
# Verify the weights are actually on disk before anything else: a wiped HF
# cache is invisible to a process that already holds the model in RAM, and
# only bites at the next (re)load — hours later, as an outage.
await _refresh_weights_state(force=True)
# Warm the model now so the very first request after a (re)start is fast,
# then keep it warm during idle gaps. Warm with the PROBE clip when we have
# one: it warms the same partial path AND establishes probe freshness, so a
# freshly-booted server isn't reported stale for its first keepwarm interval.
global _PROBE_PCM
_PROBE_PCM = _load_probe_pcm()
_probe_state["enabled"] = _PROBE_PCM is not None
try:
if _PROBE_PCM is not None:
await _run_streaming_probe()
else:
await _transcribe_buffer(bytearray(int(BYTES_PER_SEC * 0.3)))
log.info("Startup warmup complete (keepwarm interval=%.0fs, probe=%s).",
ASR_KEEPWARM_INTERVAL_SEC, _probe_state["ok"])
except Exception:
log.exception("Startup warmup failed")
asyncio.create_task(_keepwarm_loop())
asyncio.create_task(_idle_evict_loop())
log.info("Server ready.")
@app.get("/health")
async def health():
weights_present = _weights_state["present"]
# Freshness = the most recent proof the streaming path produced real text:
# a successful probe OR a successful real transcription. The probe only runs
# in idle gaps, so counting only probes would report a continuously-busy
# (and perfectly healthy) server as stale.
last_good = max(
_probe_state["last_ok_monotonic"], _last_good_transcribe_monotonic
)