diff --git a/sglang_omni/adapter/sglang_omni/models/audio8_tts/config.py b/sglang_omni/adapter/sglang_omni/models/audio8_tts/config.py --- a/sglang_omni/adapter/sglang_omni/models/audio8_tts/config.py +++ b/sglang_omni/adapter/sglang_omni/models/audio8_tts/config.py @@ -9,7 +9,6 @@ from sglang_omni.config import ( RelayConfig, StageConfig, ) -from sglang_omni.config.schema import StreamTargetConfig _PKG = "sglang_omni.models.audio8_tts.pipeline" @@ -37,7 +36,6 @@ class Audio8TTSPipelineConfig(PipelineConfig): ), get_next=f"{_PKG}.next_stage.tts_engine_next", relay=RelayConfig(device="cuda"), - stream_to=[StreamTargetConfig(to_stage="vocoder")], ), StageConfig( name="vocoder", diff --git a/sglang_omni/adapter/sglang_omni/models/audio8_tts/pipeline/stages.py b/sglang_omni/adapter/sglang_omni/models/audio8_tts/pipeline/stages.py --- a/sglang_omni/adapter/sglang_omni/models/audio8_tts/pipeline/stages.py +++ b/sglang_omni/adapter/sglang_omni/models/audio8_tts/pipeline/stages.py @@ -5,7 +5,6 @@ import base64 import importlib.util import json import logging -import os import sys from pathlib import Path from types import SimpleNamespace @@ -22,9 +21,6 @@ from sglang_omni.models.audio8_tts.pipeline.engine_io import ( build_tts_request, ) from sglang_omni.models.audio8_tts.pipeline.state_io import load_state, store_state -from sglang_omni.models.audio8_tts.pipeline.streaming_vocoder import ( - Audio8StreamingVocoderExecutor, -) from sglang_omni.models.audio8_tts.tokenizer import Audio8TokenizerAdapter, Reference from sglang_omni.proto import StagePayload @@ -239,7 +235,7 @@ def create_vocoder_executor( model_path: str, *, device: str = "cuda:0", -) -> Audio8StreamingVocoderExecutor: +) -> PreprocessingExecutor: codec = _load_codec(model_path, device) config = _load_config(model_path) warmup_codes = torch.zeros( @@ -251,13 +247,29 @@ def create_vocoder_executor( ) with torch.inference_mode(): codec.decode(warmup_codes) - return Audio8StreamingVocoderExecutor( - codec, - device=device, - eos_token_id=config.eos_token_id, - num_codebooks=config.num_codebooks, - chunk_frames=int(os.getenv("AUDIO8_TTS_STREAM_CHUNK_FRAMES", "12")), - context_frames=int(os.getenv("AUDIO8_TTS_STREAM_CONTEXT_FRAMES", "128")), - guard_frames=int(os.getenv("AUDIO8_TTS_STREAM_GUARD_FRAMES", "1")), - hop_length=config.codec_frame_size, - ) + + def vocode(payload: StagePayload) -> StagePayload: + state = load_state(payload) + if state.output_codes is None: + raise ValueError("Audio8 generation produced no codec frames") + codes = state.output_codes.to(device=device, dtype=torch.long) + with torch.inference_mode(): + audio = codec.decode(codes.unsqueeze(0))[0, 0].float().cpu() + state.audio_samples = audio + state.sample_rate = codec.sample_rate + payload = store_state(payload, state) + payload.data.update( + { + "audio_data": audio.tolist(), + "sample_rate": codec.sample_rate, + "modality": "audio", + "usage": { + "prompt_tokens": state.prompt_tokens, + "completion_tokens": state.completion_tokens, + "total_tokens": state.prompt_tokens + state.completion_tokens, + }, + } + ) + return payload + + return PreprocessingExecutor(vocode) diff --git a/sglang_omni/configs/audio8_tts_0_6b.yaml b/sglang_omni/configs/audio8_tts_0_6b.yaml --- a/sglang_omni/configs/audio8_tts_0_6b.yaml +++ b/sglang_omni/configs/audio8_tts_0_6b.yaml @@ -24,8 +24,6 @@ stages: get_next: sglang_omni.models.audio8_tts.pipeline.next_stage.tts_engine_next relay: device: cuda - stream_to: - - to_stage: vocoder - name: vocoder num_workers: 1 executor: