From b8db6ff1cebc0d6554cf0977045850bcc796b383 Mon Sep 17 00:00:00 2001 From: Kedara Studios Date: Thu, 12 Feb 2026 18:43:22 +0100 Subject: [PATCH] update README.md: list experimental branch features; GPU-resident repetition penalty and batch compaction --- README.md | 8 +- qwen_tts/core/models/modeling_qwen3_tts.py | 143 ++++++++++++++++----- qwen_tts/inference/qwen3_tts_model.py | 3 + 3 files changed, 120 insertions(+), 34 deletions(-) diff --git a/README.md b/README.md index ed96e02..3a5ad11 100644 --- a/README.md +++ b/README.md @@ -15,6 +15,12 @@ Added in this fork: - **Hann window crossfade** - click-free chunk boundaries with proper fade-in/fade-out - **Repetition penalty for streaming** - prevents token loops that cause looping audio and runaway generation. Defaults to 1.0 (disabled) because streaming generates frame-by-frame with CUDA graph constraints where repetition manifests differently than the non-streaming path (which defaults to 1.05) +Experiments on branch: [wip/experimental](https://github.com/rekuenkdr/Qwen3-TTS-streaming/tree/wip/experimental) +- **`generate_fast()` codebook predictor** - lightweight codebook generation that skips HuggingFace `generate()` overhead for the 31-step autoregressive loop (1.13x faster per-frame) +- **Manual CUDA graph capture for codebook predictor** - captures the entire 31-step codebook loop as a single CUDA graph replay (2.15x faster per-frame, 12.94ms vs 27.88ms baseline) +- **Batch streaming** - generates audio for multiple texts in parallel via `batch_stream_generate_voice_clone()`, with per-item state tracking and independent EOS detection +- **Async CUDA stream decoding** - overlaps AR token generation with speech decoding on a separate CUDA stream (disabled by default, no measurable speedup on single GPU but may show improvements in multi-GPU setups) + ## Installation ```bash @@ -178,7 +184,7 @@ model.enable_streaming_optimizations( | `use_fast_codebook` | False | Use fast codebook generation (experimental) | | `compile_codebook_predictor` | True | Apply torch.compile to codebook predictor | ---- + Based on: - [QwenLM/Qwen3-TTS](https://github.com/QwenLM/Qwen3-TTS) diff --git a/qwen_tts/core/models/modeling_qwen3_tts.py b/qwen_tts/core/models/modeling_qwen3_tts.py index 9466d0f..9c16117 100644 --- a/qwen_tts/core/models/modeling_qwen3_tts.py +++ b/qwen_tts/core/models/modeling_qwen3_tts.py @@ -2751,7 +2751,13 @@ class Qwen3TTSForConditionalGeneration(Qwen3TTSPreTrainedModel, GenerationMixin) decoded_tail: Optional[np.ndarray] = None frames_since_emit = 0 total_frames_emitted = 0 # Track how many frames we've already emitted audio for - generated_token_ids: list[int] = [token.item()] # Track first-codebook tokens for repetition penalty + + # GPU-resident circular buffer for repetition penalty + if repetition_penalty != 1.0: + rp_window = repetition_penalty_window if repetition_penalty_window > 0 else max_frames + rp_history = torch.full((1, rp_window), vocab_size, device=token.device, dtype=torch.long) + rp_history[0, 0] = token[0] + rp_step = 1 for step_idx in range(max_frames): # Mark step begin for CUDA graphs to avoid tensor overwrite errors @@ -2794,20 +2800,27 @@ class Qwen3TTSForConditionalGeneration(Qwen3TTSPreTrainedModel, GenerationMixin) # Sample next token for first codebook step_logits = step_out.logits[:, -1, :] - # Apply repetition penalty to recently generated tokens (windowed) - if repetition_penalty != 1.0 and len(generated_token_ids) > 0: - recent = generated_token_ids[-repetition_penalty_window:] if repetition_penalty_window > 0 else generated_token_ids - prev_ids = torch.tensor(list(set(recent)), device=step_logits.device) - scores = torch.gather(step_logits[0], 0, prev_ids) - scores = torch.where(scores > 0, scores / repetition_penalty, scores * repetition_penalty) - step_logits[0].scatter_(0, prev_ids, scores) + # Apply repetition penalty via GPU-resident circular buffer + if repetition_penalty != 1.0: + presence = torch.zeros(1, vocab_size + 1, device=step_logits.device, dtype=torch.bool) + presence.scatter_(1, rp_history, True) + penalty_mask = presence[:, :vocab_size] + penalized = torch.where( + step_logits > 0, + step_logits / repetition_penalty, + step_logits * repetition_penalty, + ) + step_logits = torch.where(penalty_mask, penalized, step_logits) if do_sample: token = _sample_next_token(step_logits, temperature, top_k, top_p, suppress_tokens) else: token = torch.argmax(step_logits, dim=-1) - generated_token_ids.append(token.item()) + if repetition_penalty != 1.0: + pos = rp_step % rp_window + rp_history[0, pos] = token[0] + rp_step += 1 frames_since_emit += 1 @@ -2964,6 +2977,8 @@ class Qwen3TTSForConditionalGeneration(Qwen3TTSPreTrainedModel, GenerationMixin) 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 + # Batch compaction: remove finished items from GPU tensors + compact_finished: bool = True, ) -> Generator[tuple[list[np.ndarray], int], None, None]: """ Batch streaming audio generation, yielding lists of PCM chunks as they are generated. @@ -3077,12 +3092,22 @@ class Qwen3TTSForConditionalGeneration(Qwen3TTSPreTrainedModel, GenerationMixin) 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]] = [[token[b].item()] for b in range(B)] finished: list[bool] = [False] * B sr = 24000 # default sample rate, updated on first decode - # Shared frame counter (items advance in lockstep) - frames_since_emit = 0 + # GPU-resident circular buffer for vectorized repetition penalty + if repetition_penalty != 1.0: + rp_window = repetition_penalty_window if repetition_penalty_window > 0 else max_frames + rp_history = torch.full((B, rp_window), vocab_size, device=token.device, dtype=torch.long) + rp_history[:, 0] = token # first sampled token + rp_step = 1 + + # Per-item frame counter (items advance in lockstep but counters survive compaction) + frames_since_emit = [0] * B + + # Batch compaction state: compact idx -> original batch idx + active_to_orig = list(range(B)) + B_active = B for step_idx in range(max_frames): torch.compiler.cudagraph_mark_step_begin() @@ -3111,41 +3136,83 @@ class Qwen3TTSForConditionalGeneration(Qwen3TTSPreTrainedModel, GenerationMixin) # 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]: + # Per-item EOS check and codes buffer append (using compact indices) + newly_finished_compact = set() + for ci in range(B_active): + orig = active_to_orig[ci] + if finished[orig]: continue - if codec_ids[b, 0].item() in eos_ids: - finished[b] = True + if codec_ids[ci, 0].item() in eos_ids: + finished[orig] = True + newly_finished_compact.add(ci) continue - codes_buffers[b].append(codec_ids[b].detach()) + codes_buffers[orig].append(codec_ids[ci].detach()) if all(finished): break - # Sample next token with per-item repetition penalty - step_logits = step_out.logits[:, -1, :].clone() # [B, vocab] + # --- Batch compaction: remove finished items from GPU tensors --- + compacted_this_step = False + if compact_finished and newly_finished_compact: + still_active = [ci for ci in range(B_active) if ci not in newly_finished_compact] + if still_active and len(still_active) < B_active: + compacted_this_step = True + active_idx = torch.tensor(still_active, device=token.device) + + # Slice KV cache along batch dimension (works for all cache types) + past_key_values.reorder_cache(active_idx) + + # Slice other batch-dimension tensors + past_hidden = past_hidden.index_select(0, active_idx) + trailing_text_hiddens = trailing_text_hiddens.index_select(0, active_idx) + + if self.talker.rope_deltas is not None: + self.talker.rope_deltas = self.talker.rope_deltas.index_select(0, active_idx) + + if repetition_penalty != 1.0: + rp_history = rp_history.index_select(0, active_idx) + + active_to_orig = [active_to_orig[ci] for ci in still_active] + B_active = len(active_to_orig) + + # Sample next token with vectorized repetition penalty + step_logits = step_out.logits[:, -1, :].clone() # [old_B_active, vocab] + if compacted_this_step: + step_logits = step_logits.index_select(0, active_idx) if repetition_penalty != 1.0: - for b in range(B): - if finished[b] or len(generated_token_ids[b]) == 0: - continue - recent = generated_token_ids[b][-repetition_penalty_window:] if repetition_penalty_window > 0 else generated_token_ids[b] - prev_ids = torch.tensor(list(set(recent)), 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) + presence = torch.zeros(B_active, vocab_size + 1, device=step_logits.device, dtype=torch.bool) + presence.scatter_(1, rp_history, True) + penalty_mask = presence[:, :vocab_size] + + # Zero out finished items (only needed when not compacting) + if not compact_finished: + for ci in range(B_active): + orig = active_to_orig[ci] + if finished[orig]: + penalty_mask[ci] = False + + penalized = torch.where( + step_logits > 0, + step_logits / repetition_penalty, + step_logits * repetition_penalty, + ) + step_logits = torch.where(penalty_mask, penalized, step_logits) if do_sample: token = _sample_next_token(step_logits, temperature, top_k, top_p, suppress_tokens) else: token = torch.argmax(step_logits, dim=-1) + if repetition_penalty != 1.0: + pos = rp_step % rp_window + rp_history[:, pos] = token + rp_step += 1 + + # Per-item frame counter increment for b in range(B): if not finished[b]: - generated_token_ids[b].append(token[b].item()) - - frames_since_emit += 1 + frames_since_emit[b] += 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) @@ -3164,9 +3231,16 @@ class Qwen3TTSForConditionalGeneration(Qwen3TTSPreTrainedModel, GenerationMixin) current_decode_window = decode_window_frames current_use_optimized = use_optimized_decode - if frames_since_emit < current_emit_every: + max_since_emit = max( + (frames_since_emit[b] for b in range(B) if not finished[b]), + default=0, + ) + if max_since_emit < current_emit_every: continue - frames_since_emit = 0 + # Reset all active items' counters + for b in range(B): + if not finished[b]: + frames_since_emit[b] = 0 # Decode per-item and build chunks list samples_per_frame = self.speech_tokenizer.get_decode_upsample_rate() @@ -3348,6 +3422,9 @@ class Qwen3TTSForConditionalGeneration(Qwen3TTSPreTrainedModel, GenerationMixin) if any(c.size > 0 for c in flush_chunks): yield flush_chunks, flush_sr + # Prevent stale batch-sized tensor from lingering after compaction + self.talker.rope_deltas = None + __all__ = [ "Qwen3TTSForConditionalGeneration", diff --git a/qwen_tts/inference/qwen3_tts_model.py b/qwen_tts/inference/qwen3_tts_model.py index c301a3c..ac0d03d 100644 --- a/qwen_tts/inference/qwen3_tts_model.py +++ b/qwen_tts/inference/qwen3_tts_model.py @@ -845,6 +845,8 @@ class Qwen3TTSModel: repetition_penalty_window: int = 100, # Repetition penalty (disabled by default for streaming to avoid vocabulary starvation) repetition_penalty: float = 1.0, + # Batch compaction: remove finished items from GPU tensors + compact_finished: bool = True, **kwargs, ) -> Generator[Tuple[List[np.ndarray], int], None, None]: """ @@ -957,6 +959,7 @@ class Qwen3TTSModel: first_chunk_frames=first_chunk_frames, repetition_penalty=repetition_penalty, repetition_penalty_window=repetition_penalty_window, + compact_finished=compact_finished, **gen_kwargs, ): yield chunks_list, sr