Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
104 changes: 92 additions & 12 deletions mlx_audio/tts/models/qwen3_tts/qwen3_tts.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,75 @@ def _apply_probability_filters(
return mx.where(logprobs == -mx.inf, -float("inf"), logits)


# Soft speech budget: ~3-5 codec tokens per text token at 12.5 Hz, with ~50% margin.
_EOS_EXPECTED_TOKEN_FACTOR = 6
_EOS_EXPECTED_TOKEN_FLOOR = 75
_EOS_BIAS_SCALE = 0.25
_EOS_BIAS_MAX = 20.0


def _expected_codec_token_budget(text_token_count: int) -> int:
"""Expected codec-token budget for an utterance before EOS bias ramps up."""
return max(
_EOS_EXPECTED_TOKEN_FLOOR, int(text_token_count) * _EOS_EXPECTED_TOKEN_FACTOR
)


def _progressive_eos_bias(
step: int,
expected_tokens: int,
*,
scale: float = _EOS_BIAS_SCALE,
max_bias: float = _EOS_BIAS_MAX,
) -> float:
"""Growing EOS logit bias after the expected speech budget is exceeded.

Keeps max_tokens untouched (see #695) while encouraging codec EOS once the
model starts emitting post-speech pads / breath frames.
"""
if step <= expected_tokens:
return 0.0
return min(max_bias, scale * float(step - expected_tokens))


def _add_eos_bias(
logits: mx.array,
eos_token_id: Optional[int],
eos_bias: Union[float, List[float], mx.array],
) -> mx.array:
"""Add a scalar or per-sequence bias to the EOS logit column."""
if eos_token_id is None or eos_token_id >= logits.shape[-1]:
return logits

batch = logits.shape[0]
if isinstance(eos_bias, (list, tuple)):
if len(eos_bias) != batch:
raise ValueError(
f"eos_bias length ({len(eos_bias)}) must match batch size ({batch})"
)
if not any(b != 0 for b in eos_bias):
return logits
bias = mx.array(eos_bias, dtype=logits.dtype)[:, None]
elif isinstance(eos_bias, mx.array):
if eos_bias.size == 0:
return logits
bias = eos_bias.astype(logits.dtype)
if bias.ndim == 1:
bias = bias[:, None]
if bias.shape[0] != batch:
raise ValueError(
f"eos_bias batch ({bias.shape[0]}) must match logits batch ({batch})"
)
else:
if eos_bias == 0.0:
return logits
bias = mx.full((batch, 1), float(eos_bias), dtype=logits.dtype)

eos_idx = mx.full((batch, 1), eos_token_id, dtype=mx.int32)
eos_vals = mx.take_along_axis(logits, eos_idx, axis=-1) + bias
return mx.put_along_axis(logits, eos_idx, eos_vals, axis=-1)


def mel_spectrogram(
audio: mx.array,
n_fft: int = 1024,
Expand Down Expand Up @@ -809,6 +878,7 @@ def _sample_token(
suppress_tokens: Optional[List[int]] = None,
eos_token_id: Optional[int] = None,
min_p: float = 0.0,
eos_bias: float = 0.0,
) -> mx.array:

logits = logits[:, -1, :] # Get last position [1, vocab_size]
Expand Down Expand Up @@ -843,11 +913,14 @@ def _sample_token(

# Greedy decoding if temperature is 0
if temperature <= 0:
logits = _add_eos_bias(logits, eos_token_id, eos_bias)
return mx.argmax(logits, axis=-1, keepdims=True)

if temperature != 1.0:
logits = logits / temperature

logits = _add_eos_bias(logits, eos_token_id, eos_bias)

eos_logit = None
if eos_token_id is not None and eos_token_id < logits.shape[-1]:
eos_logit = logits[:, eos_token_id : eos_token_id + 1]
Expand Down Expand Up @@ -875,6 +948,7 @@ def _sample_token_batch(
suppress_tokens: Optional[List[int]] = None,
eos_token_id: Optional[int] = None,
min_p: float = 0.0,
eos_bias: Union[float, List[float], mx.array] = 0.0,
) -> mx.array:
"""Batched sampling from [batch, seq_len, vocab] logits. Returns [batch, 1]."""

Expand Down Expand Up @@ -917,11 +991,14 @@ def _sample_token_batch(

# Greedy decoding
if temperature <= 0:
logits = _add_eos_bias(logits, eos_token_id, eos_bias)
return mx.argmax(logits, axis=-1, keepdims=True)

if temperature != 1.0:
logits = logits / temperature

logits = _add_eos_bias(logits, eos_token_id, eos_bias)

# Preserve EOS logit before filtering
eos_logit = None
if eos_token_id is not None and eos_token_id < logits.shape[-1]:
Expand Down Expand Up @@ -1823,10 +1900,10 @@ def batch_generate(
attention_mask = batch_inputs.attention_mask
mx.eval(input_embeds, trailing_text_hidden, tts_pad_embed, attention_mask)

per_seq_max_tokens = [max_tokens] * batch_size
per_seq_expected_tokens = None
if use_icl:
per_seq_max_tokens = [
min(max_tokens, max(75, len(self.tokenizer.encode(text)) * 6))
per_seq_expected_tokens = [
_expected_codec_token_budget(len(self.tokenizer.encode(text)))
for text in texts
]

Expand Down Expand Up @@ -1874,6 +1951,13 @@ def batch_generate(
attention_mask=attention_mask,
)

step_eos_bias: Union[float, List[float]] = 0.0
if per_seq_expected_tokens is not None:
step_eos_bias = [
_progressive_eos_bias(step, expected)
for expected in per_seq_expected_tokens
]

# Batched sampling — no per-sequence bool()/int() calls
sampled_tokens = self._sample_token_batch(
logits,
Expand All @@ -1884,6 +1968,7 @@ def batch_generate(
generated_tokens_per_seq=generated_token_ids,
suppress_tokens=suppress_tokens,
eos_token_id=eos_token_id,
eos_bias=step_eos_bias,
) # [batch, 1]

# Mask finished sequences to EOS (vectorized, no sync)
Expand Down Expand Up @@ -1931,15 +2016,6 @@ def batch_generate(
all_codes[b : b + 1]
) # [1, num_code_groups]

if use_icl:
finished_cpu = [
finished_cpu[b] or len(generated_codes[b]) >= per_seq_max_tokens[b]
for b in range(batch_size)
]
finished = mx.array(finished_cpu, dtype=mx.bool_)
if all(finished_cpu):
break

# Extend attention_mask by one column of 1s (skipped for bs=1)
if attention_mask is not None:
attention_mask = mx.concatenate(
Expand Down Expand Up @@ -2249,6 +2325,7 @@ def _generate_icl(
# Honor the caller-provided max_tokens; matches the base generation path.
# A text-length-derived cap could clip slow or expressive utterances.
effective_max_tokens = max_tokens
expected_tokens = _expected_codec_token_budget(len(self.tokenizer.encode(text)))

# Initialize cache
cache = self.talker.make_cache()
Expand Down Expand Up @@ -2294,6 +2371,7 @@ def _generate_icl(
generated_tokens=(generated_token_ids if generated_token_ids else None),
suppress_tokens=suppress_tokens,
eos_token_id=eos_token_id,
eos_bias=_progressive_eos_bias(step, expected_tokens),
)

# Lazy EOS check — defer sync to batch with input_embeds eval
Expand Down Expand Up @@ -2561,6 +2639,7 @@ def _generate_with_instruct(
# Honor the caller-provided max_tokens; matches the base generation path.
# A text-length-derived cap could clip slow or expressive utterances.
effective_max_tokens = max_tokens
expected_tokens = _expected_codec_token_budget(len(self.tokenizer.encode(text)))

# Initialize cache
cache = self.talker.make_cache()
Expand Down Expand Up @@ -2606,6 +2685,7 @@ def _generate_with_instruct(
generated_tokens=(generated_token_ids if generated_token_ids else None),
suppress_tokens=suppress_tokens,
eos_token_id=eos_token_id,
eos_bias=_progressive_eos_bias(step, expected_tokens),
)

# Lazy EOS check — defer sync to batch with input_embeds eval
Expand Down
85 changes: 84 additions & 1 deletion mlx_audio/tts/tests/test_qwen3_tts.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,12 @@
import mlx.core as mx
import numpy as np

from mlx_audio.tts.models.qwen3_tts.qwen3_tts import Model, mel_spectrogram
from mlx_audio.tts.models.qwen3_tts.qwen3_tts import (
Model,
_expected_codec_token_budget,
_progressive_eos_bias,
mel_spectrogram,
)
from mlx_audio.tts.models.qwen3_tts.speaker_encoder import (
TimeDelayNetBlock,
reflect_pad_1d,
Expand Down Expand Up @@ -544,5 +549,83 @@ def test_generate_icl_honors_explicit_max_tokens(self):
self.assertEqual(results[-1].token_count, 120)


class TestQwen3TTSEosBias(unittest.TestCase):
def test_expected_codec_token_budget_matches_legacy_heuristic(self):
self.assertEqual(_expected_codec_token_budget(10), 75)
self.assertEqual(_expected_codec_token_budget(20), 120)

def test_progressive_eos_bias_ramps_after_expected_budget(self):
expected = 75
self.assertEqual(_progressive_eos_bias(0, expected), 0.0)
self.assertEqual(_progressive_eos_bias(expected, expected), 0.0)
self.assertEqual(_progressive_eos_bias(expected + 4, expected), 1.0)
self.assertEqual(_progressive_eos_bias(expected + 1000, expected), 20.0)

def test_sample_token_applies_eos_bias_before_sampling(self):
logits = mx.array([[[3.0, 2.0, 1.0, 0.0]]], dtype=mx.float32)
eos_token_id = 3
captured = {}

def capture_sample(scores, temperature):
captured["scores"] = scores
captured["temperature"] = temperature
return mx.argmax(scores, axis=-1)

with patch(
"mlx_audio.tts.models.qwen3_tts.qwen3_tts.categorical_sampling",
side_effect=capture_sample,
):
Model._sample_token(
Model.__new__(Model),
logits,
temperature=1.0,
top_k=0,
top_p=1.0,
repetition_penalty=1.0,
eos_token_id=eos_token_id,
eos_bias=2.5,
)

scores = np.array(captured["scores"])
self.assertEqual(captured["temperature"], 1.0)
self.assertAlmostEqual(float(scores[0, eos_token_id]), 2.5)
np.testing.assert_allclose(scores[0, :eos_token_id], [3.0, 2.0, 1.0])

def test_sample_token_batch_applies_per_sequence_eos_bias(self):
logits = mx.array(
[
[[3.0, 2.0, 1.0, 0.0]],
[[0.0, 1.0, 2.0, 3.0]],
],
dtype=mx.float32,
)
eos_token_id = 0
captured = {}

def capture_sample(scores, temperature):
captured["scores"] = scores
captured["temperature"] = temperature
return mx.argmax(scores, axis=-1)

with patch(
"mlx_audio.tts.models.qwen3_tts.qwen3_tts.categorical_sampling",
side_effect=capture_sample,
):
Model._sample_token_batch(
Model.__new__(Model),
logits,
temperature=1.0,
top_k=0,
top_p=1.0,
repetition_penalty=1.0,
eos_token_id=eos_token_id,
eos_bias=[1.5, 0.0],
)

scores = np.array(captured["scores"])
self.assertAlmostEqual(float(scores[0, eos_token_id]), 4.5)
self.assertAlmostEqual(float(scores[1, eos_token_id]), 0.0)


if __name__ == "__main__":
unittest.main()