Decode SNAC in single pass to fix audio artifacts

SNAC's convolutional decoder has a receptive field spanning multiple
frames. Batched decode (28 tokens at a time) created boundary
discontinuities that produced muddled audio. Since pw-play waits for
the full WAV anyway, single-pass decode costs only ~2s extra and
produces clean audio.
This commit is contained in:
Ryan Malloy 2026-02-21 21:51:42 -07:00
parent 538b8a513e
commit bf0dfa7a5e

View File

@ -24,8 +24,6 @@ SAMPLE_RATE = 24000
SNAC_CODEBOOK_SIZE = 4096
TOKEN_PATTERN = re.compile(r"<custom_token_(\d+)>")
TOKENS_PER_FRAME = 7
FRAMES_PER_BATCH = 4
BATCH_SIZE = TOKENS_PER_FRAME * FRAMES_PER_BATCH # 28
ALL_VOICES = ["tara", "leah", "jess", "leo", "dan", "mia", "zac", "zoe"]
@ -149,13 +147,15 @@ class OrpheusEngine(TTSEngine):
loop = asyncio.get_running_loop()
snac = await self._get_snac()
# Stream tokens from llama-server via SSE
# Stream tokens from llama-server via SSE, then decode in one pass.
# SNAC's convolutional decoder has a receptive field spanning multiple
# frames, so batched decode creates boundary artifacts. Since pw-play
# waits for the full WAV anyway, single-pass decode costs only ~2s
# extra on a 50s generation — and produces clean audio.
t0 = time.time()
token_strings: list[str] = []
audio_chunks: list[np.ndarray] = []
total_tokens = 0
dropped_lines = 0
stream_interrupted = False
try:
async with self._client.stream(
@ -202,27 +202,16 @@ class OrpheusEngine(TTSEngine):
token_strings.append(f"<custom_token_{m}>")
total_tokens += 1
# Batch decode every 28 tokens (4 SNAC frames)
while len(token_strings) >= BATCH_SIZE:
batch = token_strings[:BATCH_SIZE]
token_strings = token_strings[BATCH_SIZE:]
chunk_audio = await loop.run_in_executor(
None, _tokens_to_audio, batch, snac
)
if chunk_audio is not None:
audio_chunks.append(chunk_audio)
except httpx.ConnectError as e:
raise RuntimeError(f"Cannot reach llama-server at {self._url}: {e}") from e
except (httpx.ReadError, httpx.RemoteProtocolError) as e:
if not audio_chunks:
if total_tokens < TOKENS_PER_FRAME:
raise RuntimeError(
f"llama-server connection lost with no audio decoded: {e}"
f"llama-server connection lost with no usable tokens: {e}"
) from e
stream_interrupted = True
print(
f" Warning: llama-server connection lost after {total_tokens} tokens. "
f"Using {len(audio_chunks)} partial chunks.",
"Decoding what we have.",
file=sys.stderr,
)
@ -234,32 +223,29 @@ class OrpheusEngine(TTSEngine):
file=sys.stderr,
)
# Flush remaining tokens (drop partial frame — at most 0.29ms lost)
if token_strings:
usable = len(token_strings) - (len(token_strings) % TOKENS_PER_FRAME)
if usable > 0:
chunk_audio = await loop.run_in_executor(
None, _tokens_to_audio, token_strings[:usable], snac
)
if chunk_audio is not None:
audio_chunks.append(chunk_audio)
# Drop partial frame at end (at most 6 tokens / 0.29ms lost)
usable = len(token_strings) - (len(token_strings) % TOKENS_PER_FRAME)
num_frames = total_tokens // TOKENS_PER_FRAME
num_frames = usable // TOKENS_PER_FRAME
tok_per_sec = total_tokens / gen_time if gen_time > 0 else 0
status = " (TRUNCATED)" if stream_interrupted else ""
print(
f" Orpheus: {total_tokens} tokens ({num_frames} frames) "
f"in {gen_time:.1f}s ({tok_per_sec:.1f} tok/s){status}",
f"in {gen_time:.1f}s ({tok_per_sec:.1f} tok/s)",
file=sys.stderr,
)
if not audio_chunks:
if usable < TOKENS_PER_FRAME:
raise RuntimeError(
f"Orpheus returned insufficient tokens ({total_tokens}). "
"Check llama-server logs."
)
audio = np.concatenate(audio_chunks)
# Single-pass SNAC decode — full context across all frames
audio = await loop.run_in_executor(
None, _tokens_to_audio, token_strings[:usable], snac
)
if audio is None:
raise RuntimeError("SNAC decoding produced no audio")
path = write_wav(audio, SAMPLE_RATE, prefix="orpheus-")
return TTSResult(