"""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)