Kleineres Modell
This commit is contained in:
357
docker/patch_qwen3_tts_runtime.py
Normal file
357
docker/patch_qwen3_tts_runtime.py
Normal file
@@ -0,0 +1,357 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Runtime patches for the vllm-omni Qwen3-TTS container.
|
||||
|
||||
- Propagate WebSocket session.config.seed into every generated speech segment.
|
||||
- Log WebSocket config and segment generation on the real 8091 endpoint.
|
||||
"""
|
||||
from pathlib import Path
|
||||
|
||||
STREAM = Path('/usr/local/lib/python3.12/dist-packages/vllm_omni/entrypoints/openai/serving_speech_stream.py')
|
||||
AUDIO = Path('/usr/local/lib/python3.12/dist-packages/vllm_omni/entrypoints/openai/protocol/audio.py')
|
||||
|
||||
|
||||
def replace_once(text: str, old: str, new: str, label: str) -> str:
|
||||
if new in text:
|
||||
return text
|
||||
if old not in text:
|
||||
raise RuntimeError(f'Patch anchor not found: {label}')
|
||||
return text.replace(old, new, 1)
|
||||
|
||||
|
||||
def patch_seed_field() -> None:
|
||||
"""Add an optional `seed` field to StreamingSpeechSessionConfig.
|
||||
|
||||
The streaming WebSocket config model in protocol/audio.py ships without a
|
||||
`seed` attribute, but the stream server reads `config.seed` (see seed
|
||||
propagation patch below). Without this field every session.config fails with
|
||||
AttributeError: 'StreamingSpeechSessionConfig' object has no attribute 'seed'.
|
||||
"""
|
||||
text = AUDIO.read_text()
|
||||
# Anchor on this class's unique docstring + first field: the task_type line
|
||||
# is not unique (it also appears in SpeechBatchItem / BatchSpeechRequest).
|
||||
new = replace_once(
|
||||
text,
|
||||
' """Configuration sent as the first WebSocket message for streaming TTS."""\n\n'
|
||||
' model: str | None = None\n',
|
||||
' """Configuration sent as the first WebSocket message for streaming TTS."""\n\n'
|
||||
' seed: int | None = None\n'
|
||||
' model: str | None = None\n',
|
||||
'StreamingSpeechSessionConfig.seed field',
|
||||
)
|
||||
if new != text:
|
||||
AUDIO.write_text(new)
|
||||
print('[runtime-patch] StreamingSpeechSessionConfig.seed field added', flush=True)
|
||||
else:
|
||||
print('[runtime-patch] StreamingSpeechSessionConfig.seed field already present', flush=True)
|
||||
|
||||
|
||||
_PIPELINE_CODE = '''
|
||||
|
||||
async def _pipelined_handle_session(self, websocket):
|
||||
"""Pipeline: N+1 synthesis overlaps with N audio transfer to client.
|
||||
|
||||
A synthesizer task submits each sentence to vllm as soon as the previous
|
||||
one is submitted (not after it is fully sent), so vllm starts work on N+1
|
||||
the moment GPU capacity is free from N. A separate sender task streams
|
||||
audio chunks to the WebSocket in order. Lookahead depth is controlled by
|
||||
QWEN3_TTS_PIPELINE_LOOKAHEAD (default 1).
|
||||
"""
|
||||
await websocket.accept()
|
||||
try:
|
||||
config = await self._receive_config(websocket)
|
||||
if config is None:
|
||||
return
|
||||
logger.info(
|
||||
"Streaming speech config: model=%s voice=%s task_type=%s language=%s"
|
||||
" seed=%s split=%s stream_audio=%s instructions=%r speaker_embedding_dim=%s",
|
||||
config.model, config.voice, config.task_type, config.language,
|
||||
config.seed, config.split_granularity, config.stream_audio, config.instructions,
|
||||
len(config.speaker_embedding) if config.speaker_embedding is not None else None,
|
||||
)
|
||||
if config.model and hasattr(self._speech_service, "_check_model"):
|
||||
error = await self._speech_service._check_model(
|
||||
OpenAICreateSpeechRequest(input="ping", model=config.model)
|
||||
)
|
||||
if error is not None:
|
||||
await self._send_error(websocket, str(error))
|
||||
return
|
||||
|
||||
boundary_re = SPLIT_CLAUSE if config.split_granularity == "clause" else SPLIT_SENTENCE
|
||||
splitter = SentenceSplitter(boundary_re=boundary_re)
|
||||
response_format = config.response_format or "wav"
|
||||
|
||||
sentence_q = asyncio.Queue()
|
||||
lookahead = int(_os.environ.get("QWEN3_TTS_PIPELINE_LOOKAHEAD", "1"))
|
||||
ready_q = asyncio.Queue(maxsize=max(1, lookahead))
|
||||
pending_rids = set() # all request IDs submitted but not yet completed
|
||||
|
||||
def _build_req(text):
|
||||
return OpenAICreateSpeechRequest(
|
||||
input=text, model=config.model, voice=config.voice,
|
||||
task_type=config.task_type, language=config.language,
|
||||
instructions=config.instructions, response_format=response_format,
|
||||
speed=config.speed, max_new_tokens=config.max_new_tokens,
|
||||
initial_codec_chunk_frames=config.initial_codec_chunk_frames,
|
||||
ref_audio=config.ref_audio, ref_text=config.ref_text,
|
||||
x_vector_only_mode=config.x_vector_only_mode,
|
||||
speaker_embedding=config.speaker_embedding,
|
||||
stream=config.stream_audio, seed=config.seed,
|
||||
)
|
||||
|
||||
async def _abort_all():
|
||||
for rid in list(pending_rids):
|
||||
try:
|
||||
await self._speech_service.engine_client.abort(rid)
|
||||
except Exception:
|
||||
pass
|
||||
pending_rids.clear()
|
||||
|
||||
async def _synthesizer():
|
||||
while True:
|
||||
item = await sentence_q.get()
|
||||
if item is None:
|
||||
await ready_q.put(None)
|
||||
return
|
||||
idx, text = item
|
||||
logger.info(
|
||||
"Streaming speech segment: index=%s chars=%s voice=%s"
|
||||
" task_type=%s language=%s seed=%s speaker_embedding_dim=%s",
|
||||
idx, len(text), config.voice, config.task_type, config.language,
|
||||
config.seed,
|
||||
len(config.speaker_embedding) if config.speaker_embedding is not None else None,
|
||||
)
|
||||
try:
|
||||
req = _build_req(text)
|
||||
if config.stream_audio:
|
||||
rid, gen, _ = await self._speech_service._prepare_speech_generation(req)
|
||||
pending_rids.add(rid)
|
||||
await ready_q.put(("stream", idx, text, rid, gen))
|
||||
else:
|
||||
audio, _ = await self._speech_service._generate_audio_bytes(req)
|
||||
await ready_q.put(("bytes", idx, text, None, audio))
|
||||
except Exception as exc:
|
||||
await ready_q.put(("error", idx, text, None, exc))
|
||||
|
||||
async def _sender():
|
||||
total = 0
|
||||
while True:
|
||||
item = await ready_q.get()
|
||||
if item is None:
|
||||
await websocket.send_json({"type": "session.done", "total_sentences": total})
|
||||
return
|
||||
kind, idx, text, extra, data = item
|
||||
if kind == "error":
|
||||
logger.error("Generation failed for sentence %d: %s", idx, data)
|
||||
await self._send_error(websocket, f"Generation failed for sentence {idx}: {data}")
|
||||
total += 1
|
||||
continue
|
||||
hdr = {"type": "audio.start", "sentence_index": idx,
|
||||
"sentence_text": text, "format": response_format}
|
||||
if config.stream_audio and response_format == "pcm":
|
||||
hdr["sample_rate"] = _PCM_SAMPLE_RATE
|
||||
await websocket.send_json(hdr)
|
||||
if config.stream_audio and response_format == "pcm" and idx == 0 and _LEADIN_SILENCE_BYTES:
|
||||
await websocket.send_bytes(_LEADIN_SILENCE_BYTES)
|
||||
sent_bytes = 0
|
||||
failed = False
|
||||
try:
|
||||
if kind == "stream":
|
||||
async with aclosing(
|
||||
self._speech_service._generate_pcm_chunks(data, extra)
|
||||
) as stream:
|
||||
async for chunk in stream:
|
||||
sent_bytes += len(chunk)
|
||||
await websocket.send_bytes(chunk)
|
||||
pending_rids.discard(extra)
|
||||
else:
|
||||
sent_bytes = len(data)
|
||||
await websocket.send_bytes(data)
|
||||
if response_format == "pcm" and _TRAILING_SILENCE_BYTES:
|
||||
await websocket.send_bytes(_TRAILING_SILENCE_BYTES)
|
||||
except WebSocketDisconnect:
|
||||
raise
|
||||
except Exception as exc:
|
||||
failed = True
|
||||
logger.error("Send failed for sentence %d: %s", idx, exc)
|
||||
await self._send_error(websocket, f"Send failed for sentence {idx}: {exc}")
|
||||
finally:
|
||||
try:
|
||||
await websocket.send_json({
|
||||
"type": "audio.done", "sentence_index": idx,
|
||||
"total_bytes": sent_bytes, "error": failed,
|
||||
})
|
||||
except Exception:
|
||||
pass
|
||||
total += 1
|
||||
|
||||
synth_task = asyncio.create_task(_synthesizer())
|
||||
send_task = asyncio.create_task(_sender())
|
||||
sentence_count = 0
|
||||
try:
|
||||
while True:
|
||||
try:
|
||||
raw = await asyncio.wait_for(
|
||||
websocket.receive_text(), timeout=self._idle_timeout
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
await self._send_error(websocket, "Idle timeout: no message received")
|
||||
synth_task.cancel()
|
||||
send_task.cancel()
|
||||
await asyncio.gather(synth_task, send_task, return_exceptions=True)
|
||||
await _abort_all()
|
||||
return
|
||||
if len(raw) > _MAX_INPUT_TEXT_MESSAGE_SIZE:
|
||||
await self._send_error(websocket, "input.text message too large")
|
||||
continue
|
||||
try:
|
||||
msg = json.loads(raw)
|
||||
except json.JSONDecodeError:
|
||||
await self._send_error(websocket, "Invalid JSON message")
|
||||
continue
|
||||
if not isinstance(msg, dict):
|
||||
await self._send_error(websocket, "WebSocket messages must be JSON objects")
|
||||
continue
|
||||
mtype = msg.get("type")
|
||||
if mtype == "input.text":
|
||||
t = msg.get("text", "")
|
||||
if not isinstance(t, str):
|
||||
await self._send_error(websocket, "input.text requires a string value")
|
||||
continue
|
||||
for s in splitter.add_text(t):
|
||||
await sentence_q.put((sentence_count, s))
|
||||
sentence_count += 1
|
||||
elif mtype == "input.done":
|
||||
rem = splitter.flush()
|
||||
if rem:
|
||||
await sentence_q.put((sentence_count, rem))
|
||||
sentence_count += 1
|
||||
await sentence_q.put(None)
|
||||
break
|
||||
else:
|
||||
await self._send_error(websocket, f"Unknown message type: {mtype}")
|
||||
await asyncio.gather(synth_task, send_task)
|
||||
except BaseException:
|
||||
synth_task.cancel()
|
||||
send_task.cancel()
|
||||
await asyncio.gather(synth_task, send_task, return_exceptions=True)
|
||||
await _abort_all()
|
||||
raise
|
||||
except WebSocketDisconnect:
|
||||
logger.info("Streaming speech: client disconnected")
|
||||
except Exception as e:
|
||||
logger.exception("Streaming speech session error: %s", e)
|
||||
try:
|
||||
await self._send_error(websocket, f"Internal error: {e}")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
OmniStreamingSpeechHandler.handle_session = _pipelined_handle_session
|
||||
'''
|
||||
|
||||
|
||||
def patch_pipeline_session() -> None:
|
||||
"""Append pipelined handle_session to serving_speech_stream.py.
|
||||
|
||||
Synthesis of sentence N+1 is submitted to vllm as soon as N's request
|
||||
is prepared, so vllm can start N+1 the moment N's GPU work finishes —
|
||||
overlapping with the network transfer of N's audio to the client.
|
||||
"""
|
||||
text = STREAM.read_text()
|
||||
sentinel = 'OmniStreamingSpeechHandler.handle_session = _pipelined_handle_session'
|
||||
if sentinel in text:
|
||||
print('[runtime-patch] pipeline session handler already present', flush=True)
|
||||
return
|
||||
STREAM.write_text(text + _PIPELINE_CODE)
|
||||
print('[runtime-patch] pipeline session handler injected', flush=True)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
patch_seed_field()
|
||||
|
||||
text = STREAM.read_text()
|
||||
original = text
|
||||
|
||||
text = replace_once(
|
||||
text,
|
||||
' if config.model and hasattr(self._speech_service, "_check_model"):\n',
|
||||
' logger.info(\n'
|
||||
' "Streaming speech config: model=%s voice=%s task_type=%s language=%s seed=%s split=%s stream_audio=%s instructions=%r speaker_embedding_dim=%s",\n'
|
||||
' config.model,\n'
|
||||
' config.voice,\n'
|
||||
' config.task_type,\n'
|
||||
' config.language,\n'
|
||||
' config.seed,\n'
|
||||
' config.split_granularity,\n'
|
||||
' config.stream_audio,\n'
|
||||
' config.instructions,\n'
|
||||
' len(config.speaker_embedding) if config.speaker_embedding is not None else None,\n'
|
||||
' )\n\n'
|
||||
' if config.model and hasattr(self._speech_service, "_check_model"):\n',
|
||||
'config logging',
|
||||
)
|
||||
|
||||
text = replace_once(
|
||||
text,
|
||||
' stream=config.stream_audio,\n )\n',
|
||||
' stream=config.stream_audio,\n seed=config.seed,\n )\n'
|
||||
' logger.info(\n'
|
||||
' "Streaming speech segment: index=%s chars=%s voice=%s task_type=%s language=%s seed=%s speaker_embedding_dim=%s",\n'
|
||||
' sentence_index,\n'
|
||||
' len(sentence_text),\n'
|
||||
' config.voice,\n'
|
||||
' config.task_type,\n'
|
||||
' config.language,\n'
|
||||
' config.seed,\n'
|
||||
' len(config.speaker_embedding) if config.speaker_embedding is not None else None,\n'
|
||||
' )\n',
|
||||
'seed propagation and segment logging',
|
||||
)
|
||||
|
||||
# --- Silence padding (PCM streaming only) ---
|
||||
# Lead-in silence at session start gives the client time to set up smooth
|
||||
# playback before real audio (mitigates start-of-stream stutter). Trailing
|
||||
# silence avoids the abrupt cut-off at the end. Both in ms via env vars.
|
||||
text = replace_once(
|
||||
text,
|
||||
'_PCM_SAMPLE_RATE = 24000\n',
|
||||
'_PCM_SAMPLE_RATE = 24000\n'
|
||||
'import os as _os\n'
|
||||
'_LEADIN_SILENCE_MS = int(_os.environ.get("QWEN3_TTS_LEADIN_SILENCE_MS", "250"))\n'
|
||||
'_TRAILING_SILENCE_MS = int(_os.environ.get("QWEN3_TTS_TRAILING_SILENCE_MS", "300"))\n'
|
||||
'_LEADIN_SILENCE_BYTES = b"\\x00" * (int(_PCM_SAMPLE_RATE * _LEADIN_SILENCE_MS / 1000) * 2)\n'
|
||||
'_TRAILING_SILENCE_BYTES = b"\\x00" * (int(_PCM_SAMPLE_RATE * _TRAILING_SILENCE_MS / 1000) * 2)\n',
|
||||
'silence padding constants',
|
||||
)
|
||||
|
||||
text = replace_once(
|
||||
text,
|
||||
' await websocket.send_json(start_payload)\n',
|
||||
' await websocket.send_json(start_payload)\n'
|
||||
' if config.stream_audio and response_format == "pcm" and sentence_index == 0 and _LEADIN_SILENCE_BYTES:\n'
|
||||
' await websocket.send_bytes(_LEADIN_SILENCE_BYTES)\n',
|
||||
'lead-in silence',
|
||||
)
|
||||
|
||||
text = replace_once(
|
||||
text,
|
||||
' # Send session.done\n'
|
||||
' await websocket.send_json(\n',
|
||||
' if config.stream_audio and config.response_format == "pcm" and _TRAILING_SILENCE_BYTES:\n'
|
||||
' await websocket.send_bytes(_TRAILING_SILENCE_BYTES)\n\n'
|
||||
' # Send session.done\n'
|
||||
' await websocket.send_json(\n',
|
||||
'trailing silence',
|
||||
)
|
||||
|
||||
if text != original:
|
||||
STREAM.write_text(text)
|
||||
print('[runtime-patch] WebSocket seed propagation/logging + silence padding applied', flush=True)
|
||||
else:
|
||||
print('[runtime-patch] WebSocket seed propagation/logging already present', flush=True)
|
||||
|
||||
patch_pipeline_session()
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
Reference in New Issue
Block a user