Merge pull request #8 from rekuenkdr/feature/batch-streaming

feat: add batch streaming generation for parallel multi-item TTS
This commit is contained in:
rekuenkdr
2026-02-08 17:30:24 +01:00
committed by GitHub
3 changed files with 599 additions and 0 deletions
+132
View File
@@ -0,0 +1,132 @@
"""Test batch streaming voice clone generation.
Generates audio for multiple texts (potentially with different voices) in a
single batched pass through the transformer. All items advance in lockstep.
"""
import time
import numpy as np
import soundfile as sf
import torch
from qwen_tts import Qwen3TTSModel
def log_time(start, operation):
elapsed = time.time() - start
print(f"[{elapsed:.2f}s] {operation}")
return time.time()
total_start = time.time()
# Load model
start = time.time()
model = Qwen3TTSModel.from_pretrained(
"Qwen/Qwen3-TTS-12Hz-1.7B-Base",
device_map="cuda:0",
dtype=torch.bfloat16,
attn_implementation="flash_attention_2",
)
start = log_time(start, "Load Base model")
# Create a voice clone prompt (using the same voice for all items in this test)
ref_audio_path = "kuklina-1.wav" # Replace with your reference audio
ref_text = (
"Это брат Кэти, моей одноклассницы. А что у тебя с рукой? И почему ты голая? "
"У него ведь куча наград по боевым искусствам."
)
voice_prompt = model.create_voice_clone_prompt(
ref_audio=ref_audio_path,
ref_text=ref_text,
)
start = log_time(start, "Create voice clone prompt")
# Batch items: different texts, same voice (broadcast)
texts = [
"Hello! This is the first batch item with a short sentence.",
"And this is the second batch item. It has a bit more text to synthesize.",
"Third item here. Testing batch streaming with multiple items at once!",
]
languages = ["English", "English", "English"]
# ============== Batch Streaming ==============
print(f"\n--- Batch streaming ({len(texts)} items) ---")
start = time.time()
# Accumulate per-item chunks
item_chunks: list[list[np.ndarray]] = [[] for _ in range(len(texts))]
chunk_count = 0
sr = 24000
for chunks_list, chunk_sr in model.batch_stream_generate_voice_clone(
text=texts,
language=languages,
voice_clone_prompt=voice_prompt,
emit_every_frames=8,
decode_window_frames=80,
overlap_samples=512,
max_frames=8000,
first_chunk_emit_every=5,
first_chunk_decode_window=48,
first_chunk_frames=48,
):
sr = chunk_sr
chunk_count += 1
sizes = []
for b, chunk in enumerate(chunks_list):
if chunk.size > 0:
item_chunks[b].append(chunk)
sizes.append(f"item{b}={chunk.size}")
else:
sizes.append(f"item{b}=empty")
if chunk_count <= 5 or chunk_count % 10 == 0:
print(f" Chunk {chunk_count}: {', '.join(sizes)}")
start = log_time(start, f"Batch streaming done ({chunk_count} chunks)")
# Save per-item outputs
for i, chunks in enumerate(item_chunks):
if chunks:
combined = np.concatenate(chunks)
filename = f"batch_item_{i}.wav"
sf.write(filename, combined, sr)
duration_ms = len(combined) / sr * 1000
print(f" Saved {filename}: {duration_ms:.0f}ms, {len(combined)} samples")
else:
print(f" Item {i}: no audio generated")
# ============== Compare: Sequential single-item streaming ==============
print(f"\n--- Sequential single-item streaming ({len(texts)} items) ---")
start = time.time()
for i, text in enumerate(texts):
item_single_chunks = []
for chunk, chunk_sr in model.stream_generate_voice_clone(
text=text,
language=languages[i],
voice_clone_prompt=voice_prompt,
emit_every_frames=8,
decode_window_frames=80,
overlap_samples=512,
max_frames=8000,
first_chunk_emit_every=5,
first_chunk_decode_window=48,
first_chunk_frames=48,
):
if chunk.size > 0:
item_single_chunks.append(chunk)
if item_single_chunks:
combined = np.concatenate(item_single_chunks)
filename = f"single_item_{i}.wav"
sf.write(filename, combined, chunk_sr)
duration_ms = len(combined) / chunk_sr * 1000
print(f" Saved {filename}: {duration_ms:.0f}ms")
start = log_time(start, "Sequential streaming done")
total_elapsed = time.time() - total_start
print(f"\nTotal time: {total_elapsed:.2f}s")
+338
View File
@@ -2922,6 +2922,344 @@ class Qwen3TTSForConditionalGeneration(Qwen3TTSPreTrainedModel, GenerationMixin)
yield wav, sr
@torch.inference_mode()
def batch_stream_generate_pcm(
self,
input_ids: list[torch.Tensor],
instruct_ids: Optional[list[torch.Tensor]] = None,
ref_ids: Optional[list[torch.Tensor]] = None,
voice_clone_prompt: Optional[list[dict]] = None,
languages: Optional[list[str]] = None,
speakers: Optional[list[str]] = None,
non_streaming_mode: bool = False,
# Sampling parameters for first codebook
do_sample: bool = True,
top_k: int = 50,
top_p: float = 1.0,
temperature: float = 0.9,
# Sub-talker parameters (for remaining code groups)
subtalker_dosample: bool = True,
subtalker_top_k: int = 50,
subtalker_top_p: float = 1.0,
subtalker_temperature: float = 0.9,
# Repetition penalty
repetition_penalty: float = 1.0,
# Streaming control
emit_every_frames: int = 8,
decode_window_frames: int = 80,
overlap_samples: int = 512,
max_frames: int = 10000,
# Optimization flags
use_optimized_decode: bool = True,
# Two-phase streaming: aggressive first chunk
first_chunk_emit_every: int = 0, # 0 = disabled, use emit_every_frames throughout
first_chunk_decode_window: int = 48,
first_chunk_frames: int = 48, # Switch to stable after this many frames
) -> Generator[tuple[list[np.ndarray], int], None, None]:
"""
Batch streaming audio generation, yielding lists of PCM chunks as they are generated.
All batch items advance in lockstep through the transformer. Per-item state is
maintained for codes buffers, decoded tails, repetition penalty tracking,
ref_code contexts, and EOS detection.
Args:
input_ids: List of input token tensors (one per batch item)
instruct_ids: Optional instruction token tensors
ref_ids: Optional reference token tensors
voice_clone_prompt: Optional voice cloning prompt dict (lists indexed per item)
languages: List of language strings (one per batch item)
speakers: Optional list of speaker names
non_streaming_mode: Whether to use non-streaming text mode
do_sample: Whether to sample (vs greedy) for first codebook
top_k: Top-k filtering for sampling
top_p: Top-p (nucleus) filtering for sampling
temperature: Sampling temperature
subtalker_*: Parameters for sub-codebook prediction
repetition_penalty: Penalty to reduce repeated tokens/codes
emit_every_frames: Emit PCM chunk every N codec frames (phase 2)
decode_window_frames: Window size for decoding (phase 2)
overlap_samples: Overlap samples for crossfade between chunks
max_frames: Maximum number of codec frames to generate
use_optimized_decode: Use CUDA graph optimized decode when available
first_chunk_emit_every: Emit interval for first chunk phase (0 = disabled)
first_chunk_decode_window: Decode window size for first chunk phase
first_chunk_frames: Switch to stable settings after this many frames
Yields:
tuple[list[np.ndarray], int]: (chunks_list, sample_rate) where chunks_list[b]
is the PCM chunk for batch item b. Finished items get empty arrays.
"""
B = len(input_ids)
# Build talker inputs (already handles batching with padding)
talker_input_embeds, talker_attention_mask, trailing_text_hiddens, tts_pad_embed = \
self._build_talker_inputs(
input_ids=input_ids,
instruct_ids=instruct_ids,
ref_ids=ref_ids,
voice_clone_prompt=voice_clone_prompt,
languages=languages,
speakers=speakers,
non_streaming_mode=non_streaming_mode,
)
# Multiple EOS tokens that can terminate generation
eos_ids = {
self.config.talker_config.codec_eos_token_id,
2150, 2157, 151670,
self.config.tts_eos_token_id,
self.config.im_end_token_id,
151643,
}
vocab_size = self.config.talker_config.vocab_size
suppress_tokens = [
i for i in range(vocab_size - 1024, vocab_size)
if i not in eos_ids
]
# Mark step begin for CUDA graphs
torch.compiler.cudagraph_mark_step_begin()
# Prefill: single batched forward pass to initialize KV cache
out = self.talker.forward(
inputs_embeds=talker_input_embeds,
attention_mask=talker_attention_mask,
use_cache=True,
output_hidden_states=True,
return_dict=True,
trailing_text_hidden=trailing_text_hiddens,
tts_pad_embed=tts_pad_embed,
generation_step=None,
past_hidden=None,
past_key_values=None,
subtalker_dosample=subtalker_dosample,
subtalker_top_k=subtalker_top_k,
subtalker_top_p=subtalker_top_p,
subtalker_temperature=subtalker_temperature,
)
past_key_values = out.past_key_values
past_hidden = out.past_hidden
generation_step = out.generation_step
# Sample first token from prefill logits [B, vocab]
last_logits = out.logits[:, -1, :]
if do_sample:
token = _sample_next_token(last_logits, temperature, top_k, top_p, suppress_tokens)
else:
token = torch.argmax(last_logits, dim=-1)
# Extract per-item ref_code context (if in ICL mode)
ref_code_contexts: list[Optional[torch.Tensor]] = [None] * B
ref_code_frames_list: list[int] = [0] * B
if voice_clone_prompt is not None:
ref_code_list = voice_clone_prompt.get("ref_code", None)
icl_mode_list = voice_clone_prompt.get("icl_mode", None)
if ref_code_list is not None and icl_mode_list is not None:
for b in range(B):
if b < len(ref_code_list) and ref_code_list[b] is not None and icl_mode_list[b]:
ref_code_contexts[b] = ref_code_list[b].to(self.talker.device)
ref_code_frames_list[b] = ref_code_contexts[b].shape[0]
# Per-item decode state
codes_buffers: list[list[torch.Tensor]] = [[] for _ in range(B)]
decoded_tails: list[Optional[np.ndarray]] = [None] * B
total_frames_emitted: list[int] = [0] * B
generated_token_ids: list[list[int]] = [[] for _ in range(B)]
finished: list[bool] = [False] * B
# Shared frame counter (items advance in lockstep)
frames_since_emit = 0
for step_idx in range(max_frames):
torch.compiler.cudagraph_mark_step_begin()
# Single-step batched forward
step_out = self.talker.forward(
input_ids=token.unsqueeze(1),
use_cache=True,
return_dict=True,
output_hidden_states=False,
past_key_values=past_key_values,
past_hidden=past_hidden,
generation_step=generation_step,
trailing_text_hidden=trailing_text_hiddens,
tts_pad_embed=tts_pad_embed,
subtalker_dosample=subtalker_dosample,
subtalker_top_k=subtalker_top_k,
subtalker_top_p=subtalker_top_p,
subtalker_temperature=subtalker_temperature,
)
past_key_values = step_out.past_key_values
past_hidden = step_out.past_hidden
generation_step = step_out.generation_step
# Get codec_ids [B, num_code_groups]
codec_ids = step_out.hidden_states[1]
# Per-item EOS check and codes buffer append
for b in range(B):
if finished[b]:
continue
if codec_ids[b, 0].item() in eos_ids:
finished[b] = True
continue
codes_buffers[b].append(codec_ids[b].detach())
if all(finished):
break
# Sample next token with per-item repetition penalty
step_logits = step_out.logits[:, -1, :].clone() # [B, vocab]
if repetition_penalty != 1.0:
for b in range(B):
if finished[b] or len(generated_token_ids[b]) == 0:
continue
prev_ids = torch.tensor(list(set(generated_token_ids[b])), device=step_logits.device)
scores = torch.gather(step_logits[b], 0, prev_ids)
scores = torch.where(scores > 0, scores / repetition_penalty, scores * repetition_penalty)
step_logits[b].scatter_(0, prev_ids, scores)
if do_sample:
token = _sample_next_token(step_logits, temperature, top_k, top_p, suppress_tokens)
else:
token = torch.argmax(step_logits, dim=-1)
for b in range(B):
if not finished[b]:
generated_token_ids[b].append(token[b].item())
frames_since_emit += 1
# Two-phase streaming: use any active item's buffer length for phase detection
# (all active items have same buffer length since they advance in lockstep)
any_active_frames = 0
for b in range(B):
if not finished[b] and len(codes_buffers[b]) > 0:
any_active_frames = len(codes_buffers[b])
break
if first_chunk_emit_every > 0 and any_active_frames < first_chunk_frames:
current_emit_every = first_chunk_emit_every
current_decode_window = first_chunk_decode_window
current_use_optimized = False
else:
current_emit_every = emit_every_frames
current_decode_window = decode_window_frames
current_use_optimized = use_optimized_decode
if frames_since_emit < current_emit_every:
continue
frames_since_emit = 0
# Decode per-item and build chunks list
samples_per_frame = self.speech_tokenizer.get_decode_upsample_rate()
step_samples = samples_per_frame * current_emit_every
blend_samples = overlap_samples
chunks_list: list[np.ndarray] = []
for b in range(B):
if finished[b] or len(codes_buffers[b]) == 0:
chunks_list.append(np.array([], dtype=np.float32))
continue
# Decode window for this item
start = max(0, len(codes_buffers[b]) - current_decode_window)
window_codes = torch.stack(codes_buffers[b][start:], dim=0)
# Add per-item ref_code context
window, _ = _add_ref_code_context(
window_codes, ref_code_contexts[b], ref_code_frames_list[b], current_decode_window
)
# Decode (each item independently through compiled decoder)
if current_use_optimized and hasattr(self.speech_tokenizer, 'decode_streaming'):
wavs, sr = self.speech_tokenizer.decode_streaming(
window.to(self.talker.device),
use_optimized=True,
pad_to_size=decode_window_frames,
)
else:
wavs, sr = self.speech_tokenizer.decode([{"audio_codes": window.to(self.talker.device)}])
wav = wavs[0].astype(np.float32)
chunk = wav[-step_samples:] if step_samples > 0 else wav
# Crossfade with previous chunk tail
if decoded_tails[b] is not None:
ov = min(blend_samples, len(decoded_tails[b]), len(chunk))
if ov > 0:
head = _crossfade(decoded_tails[b][-ov:], chunk[:ov])
chunk = np.concatenate([head, chunk[ov:]], axis=0)
# Hann fade-in on very first chunk
if decoded_tails[b] is None:
fade_len = min(blend_samples, len(chunk))
if fade_len > 0:
t = np.arange(fade_len, dtype=np.float32) / max(fade_len - 1, 1)
fade_in = 0.5 * (1 - np.cos(np.pi * t))
chunk[:fade_len] *= fade_in
decoded_tails[b] = chunk.copy()
if len(chunk) > blend_samples * 2:
chunk = chunk[:-blend_samples]
total_frames_emitted[b] = len(codes_buffers[b])
chunks_list.append(chunk)
yield chunks_list, sr
# Flush: decode remaining per-item frames
flush_chunks: list[np.ndarray] = []
flush_sr = 24000 # default
for b in range(B):
remaining_frames = len(codes_buffers[b]) - total_frames_emitted[b]
if remaining_frames <= 0:
flush_chunks.append(np.array([], dtype=np.float32))
continue
context_frames = min(total_frames_emitted[b], decode_window_frames - remaining_frames)
start_idx = total_frames_emitted[b] - context_frames
window_codes = torch.stack(codes_buffers[b][start_idx:], dim=0)
window, flush_ref_prefix_frames = _add_ref_code_context(
window_codes, ref_code_contexts[b], ref_code_frames_list[b], decode_window_frames
)
wavs, flush_sr = self.speech_tokenizer.decode([{"audio_codes": window.to(self.talker.device)}])
wav = wavs[0].astype(np.float32)
skip_frames = flush_ref_prefix_frames + context_frames
if skip_frames > 0:
samples_per_frame = len(wav) / window.shape[0]
skip_samples = int(skip_frames * samples_per_frame)
wav = wav[skip_samples:]
blend_samples = overlap_samples
if decoded_tails[b] is not None and len(wav) > 0:
ov = min(blend_samples, len(decoded_tails[b]), len(wav))
if ov > 0:
head = _crossfade(decoded_tails[b][-ov:], wav[:ov])
wav = np.concatenate([head, wav[ov:]], axis=0)
if len(wav) > blend_samples:
fade_len = min(blend_samples, len(wav))
t = np.arange(fade_len, dtype=np.float32) / max(fade_len - 1, 1)
fade_out = 0.5 * (1 + np.cos(np.pi * t))
wav[-fade_len:] *= fade_out
flush_chunks.append(wav)
if any(c.size > 0 for c in flush_chunks):
yield flush_chunks, flush_sr
__all__ = [
"Qwen3TTSForConditionalGeneration",
"Qwen3TTSTalkerForConditionalGeneration",
+129
View File
@@ -813,6 +813,135 @@ class Qwen3TTSModel:
):
yield chunk, sr
@torch.inference_mode()
def batch_stream_generate_voice_clone(
self,
text: List[str],
language: Union[str, List[str], None] = None,
voice_clone_prompt: Union[List[VoiceClonePromptItem], VoiceClonePromptItem, None] = None,
non_streaming_mode: bool = False,
# Streaming control
emit_every_frames: int = 8,
decode_window_frames: int = 80,
overlap_samples: int = 512,
max_frames: int = 10000,
# Optimization
use_optimized_decode: bool = True,
# Two-phase streaming
first_chunk_emit_every: int = 0,
first_chunk_decode_window: int = 48,
first_chunk_frames: int = 48,
**kwargs,
) -> Generator[Tuple[List[np.ndarray], int], None, None]:
"""
Batch streaming voice clone speech generation.
All batch items advance in lockstep through the transformer.
Per-item state is maintained for codes, crossfade, repetition penalty,
and ref_code context.
Args:
text: List of texts to synthesize (one per batch item).
language: Language(s) for synthesis. If str or None, broadcast to all items.
voice_clone_prompt: Pre-built VoiceClonePromptItem(s). Single item is broadcast.
non_streaming_mode: Whether to use non-streaming text input mode.
emit_every_frames: Emit interval for phase 2.
decode_window_frames: Decode window for phase 2.
overlap_samples: Overlap samples for crossfade.
max_frames: Maximum codec frames to generate.
use_optimized_decode: Use CUDA graph optimized decode.
first_chunk_emit_every: Emit interval for phase 1 (0 = disabled).
first_chunk_decode_window: Decode window for phase 1.
first_chunk_frames: Switch to phase 2 after this many frames.
**kwargs: Generation parameters (do_sample, top_k, top_p, temperature, etc.)
Yields:
Tuple[List[np.ndarray], int]: (chunks_list, sample_rate)
"""
if self.model.tts_model_type != "base":
raise ValueError(
f"model with tts_model_type={self.model.tts_model_type} "
"does not support batch_stream_generate_voice_clone"
)
if not isinstance(text, list) or len(text) < 1:
raise ValueError("text must be a non-empty list of strings")
B = len(text)
# Broadcast language
if language is None:
languages = ["Auto"] * B
elif isinstance(language, str):
languages = [language] * B
else:
languages = list(language)
if len(languages) == 1 and B > 1:
languages = languages * B
if len(languages) != B:
raise ValueError(f"Batch size mismatch: text={B}, language={len(languages)}")
self._validate_languages(languages)
# Broadcast voice_clone_prompt
if voice_clone_prompt is None:
raise ValueError("voice_clone_prompt is required for batch_stream_generate_voice_clone")
if isinstance(voice_clone_prompt, VoiceClonePromptItem):
prompt_items = [voice_clone_prompt] * B
elif isinstance(voice_clone_prompt, list):
prompt_items = voice_clone_prompt
if len(prompt_items) == 1 and B > 1:
prompt_items = prompt_items * B
if len(prompt_items) != B:
raise ValueError(f"Batch size mismatch: prompt={len(prompt_items)}, text={B}")
else:
raise ValueError("voice_clone_prompt must be VoiceClonePromptItem or list of them")
voice_clone_prompt_dict = self._prompt_items_to_voice_clone_prompt(prompt_items)
ref_texts_for_ids = [it.ref_text for it in prompt_items]
# Tokenize each text
input_texts = [self._build_assistant_text(t) for t in text]
input_ids = self._tokenize_texts(input_texts)
# Build per-item ref_ids
ref_ids = None
if ref_texts_for_ids is not None:
ref_ids = []
for rt in ref_texts_for_ids:
if rt is None or rt == "":
ref_ids.append(None)
else:
ref_tok = self._tokenize_texts([self._build_ref_text(rt)])[0]
ref_ids.append(ref_tok)
# Filter to supported generation params
gen_kwargs = self._merge_generate_kwargs(**kwargs)
supported_params = {
"do_sample", "top_k", "top_p", "temperature",
"subtalker_dosample", "subtalker_top_k", "subtalker_top_p", "subtalker_temperature",
"repetition_penalty"
}
gen_kwargs = {k: v for k, v in gen_kwargs.items() if k in supported_params}
# Delegate to batch_stream_generate_pcm
for chunks_list, sr in self.model.batch_stream_generate_pcm(
input_ids=input_ids,
ref_ids=ref_ids,
voice_clone_prompt=voice_clone_prompt_dict,
languages=languages,
non_streaming_mode=non_streaming_mode,
emit_every_frames=emit_every_frames,
decode_window_frames=decode_window_frames,
overlap_samples=overlap_samples,
max_frames=max_frames,
use_optimized_decode=use_optimized_decode,
first_chunk_emit_every=first_chunk_emit_every,
first_chunk_decode_window=first_chunk_decode_window,
first_chunk_frames=first_chunk_frames,
**gen_kwargs,
):
yield chunks_list, sr
# voice design model
@torch.no_grad()
def generate_voice_design(