Download dude_tts/model.py from penguinfish1688/duplexdataengine: direct link, hf CLI and curl.
- Browser
- Download file 7.33 kB
-
https://huggingface.co/penguinfish1688/duplexdataengine/resolve/main/dude_tts/model.py
- Command line
-
hf download hf://penguinfish1688/duplexdataengine/dude_tts/model.py
-
curl -L -o model.py https://huggingface.co/penguinfish1688/duplexdataengine/resolve/main/dude_tts/model.py
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 | |
| 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) | |
| 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) | |