penguinfish1688's picture
Release standalone duplex TTS with XML interleaver, voice presets and examples
4123b95 verified
Raw History Blame Contribute Delete
7.33 kB
"""Shared duplex TTS model: inference only."""
import torch
from torch import nn
from transformers import RepetitionPenaltyLogitsProcessor, TopKLogitsWarper, TopPLogitsWarper
from .layout import native_prefix
def select_primary(logits, history, *, eos, frame, do_sample, temperature=0.9,
top_k=50, top_p=1.0, repetition_penalty=1.05):
"""Official Qwen sampling rules, with repetition history local to one lane."""
if temperature <= 0 or repetition_penalty <= 0 or top_k < 0 or not 0 < top_p <= 1:
raise ValueError("Invalid generation settings")
scores = logits.float().clone()
eos_score = scores[:, eos].clone()
scores[:, 2048:] = -float("inf")
scores[:, eos] = eos_score
ids = torch.tensor([history], device=scores.device, dtype=torch.long)
if history and repetition_penalty != 1.0:
scores = RepetitionPenaltyLogitsProcessor(float(repetition_penalty))(ids, scores)
if frame < 2:
scores[:, eos] = -float("inf")
if not do_sample:
return scores.argmax(-1, keepdim=True)
scores = scores / temperature
if top_k:
scores = TopKLogitsWarper(top_k)(ids, scores)
if top_p < 1.0:
scores = TopPLogitsWarper(top_p)(ids, scores)
return torch.multinomial(scores.softmax(-1), 1)
class DualChannelTTS(nn.Module):
def __init__(self, wrapper, channel_init_std=0.02):
super().__init__()
self.talker = wrapper.model.talker
self.config = wrapper.model.config
self.depth_checkpointing = False
self.channel_embedding = nn.Embedding(2, self.talker.config.hidden_size, device=next(self.talker.parameters()).device, dtype=next(self.talker.parameters()).dtype)
self.channel_init_std = channel_init_std
nn.init.normal_(self.channel_embedding.weight, std=channel_init_std)
def embed(self, text_ids, codec_ids, add_channels=True, offset=0, voice_positions=None, voice_embeddings=None):
text = self.talker.text_projection(self.talker.get_text_embeddings()(text_ids.clamp_min(0)))
embeds = text * (text_ids >= 0).unsqueeze(-1)
for i in range(16):
table = self.talker.get_input_embeddings() if i == 0 else self.talker.code_predictor.get_input_embeddings()[i - 1]
ids = codec_ids[..., i]
embeds = embeds + table(ids.clamp_min(0)) * (ids >= 0).unsqueeze(-1)
if voice_positions is not None:
if voice_embeddings is None:
raise ValueError('Speaker conditioning positions require reference vectors')
batch = torch.arange(len(text_ids), device=text_ids.device)[:, None]
embeds[batch, voice_positions] = embeds[batch, voice_positions] + voice_embeddings.to(embeds.dtype)
if add_channels:
channels = (torch.arange(text_ids.shape[1], device=text_ids.device) + offset) % 2
embeds = embeds + self.channel_embedding(channels)[None]
return embeds
@torch.inference_mode()
def generate_duplex(self, record, encode_text, encode_event=None, **kwargs):
from .frontend import duplex_prefix
conditioning = duplex_prefix(record, self.config, encode_text, encode_event)
return self.generate_pair([], _conditioning=conditioning, **kwargs)
@torch.inference_mode()
def generate_pair(self, records, max_frames=512, do_sample=True, temperature=0.9, top_k=50, top_p=1.0, repetition_penalty=1.05, subtalker_dosample=None, subtalker_temperature=0.9, subtalker_top_k=50, subtalker_top_p=1.0, _conditioning=None):
"""Free generation: two independent cursors/EOS flags and one shared KV cache."""
self.eval()
device = next(self.parameters()).device
voices = {}
if _conditioning is None:
prefixes = [native_prefix(record, self.config) for record in records]
p = max((len(text) for text, _ in prefixes))
text = torch.full((1, 2 * p), -1, dtype=torch.long, device=device)
codec = torch.full((1, 2 * p, 16), -1, dtype=torch.long, device=device)
mask = torch.zeros((1, 2 * p), dtype=torch.bool, device=device)
for ch, (pt, pc) in enumerate(prefixes):
pos = 2 * torch.arange(p - len(pt), p, device=device) + ch
text[0, pos], codec[0, pos], mask[0, pos] = (pt.to(device), pc.to(device), True)
else:
text, codec, mask = [_conditioning[key][None].to(device) for key in ['text_ids', 'codec_ids', 'attention_mask']]
p = text.shape[1] // 2
voices = {key: _conditioning[key][None].to(device) for key in ['voice_positions', 'voice_embeddings']}
out = self.talker.model(inputs_embeds=self.embed(text, codec, **voices), attention_mask=mask, position_ids=torch.arange(2 * p, device=device)[None], use_cache=True, return_dict=True)
hidden, cache = (out.last_hidden_state[:, -2:], out.past_key_values)
results, done = ([[], []], [False, False])
history = [[], []]
if subtalker_dosample is None:
subtalker_dosample = do_sample
eos = self.config.talker_config.codec_eos_token_id
for frame in range(max_frames):
inputs = torch.full((1, 2, 16), -1, dtype=torch.long, device=device)
inputs[..., 0] = self.config.talker_config.codec_pad_id
next_text = torch.full((1, 2), self.config.tts_pad_token_id, dtype=torch.long, device=device)
new_mask = torch.zeros((1, 2), dtype=torch.bool, device=device)
for ch in range(2):
if done[ch]:
continue
h = hidden[:, ch:ch + 1]
logits = self.talker.codec_head(h[:, 0]).float()
first = select_primary(logits, history[ch], eos=eos, frame=frame, do_sample=do_sample, temperature=temperature, top_k=top_k, top_p=top_p, repetition_penalty=repetition_penalty)
if first.item() == eos:
done[ch] = True
continue
history[ch].append(first.item())
residual = self.talker.code_predictor.generate(inputs_embeds=torch.cat([h, self.talker.get_input_embeddings()(first)], dim=1), max_new_tokens=15, do_sample=subtalker_dosample, return_dict_in_generate=False, **{'temperature': subtalker_temperature, 'top_k': subtalker_top_k, 'top_p': subtalker_top_p} if subtalker_dosample else {})
codes = torch.cat([first, residual], dim=-1)[0]
if codes.shape != (16,):
raise RuntimeError(f'Unexpected residual decoder shape: {codes.shape}')
results[ch].append(codes)
inputs[0, ch] = codes
next_text[0, ch] = self.config.tts_pad_token_id
new_mask[0, ch] = True
if all(done):
break
start = mask.shape[1]
mask = torch.cat([mask, new_mask], dim=1)
out = self.talker.model(inputs_embeds=self.embed(next_text, inputs, offset=start), attention_mask=mask, position_ids=torch.arange(start, start + 2, device=device)[None], past_key_values=cache, use_cache=True, return_dict=True)
hidden, cache = (out.last_hidden_state, out.past_key_values)
arrays = [torch.stack(r).cpu().numpy() if r else None for r in results]
return (arrays, done)