duplexdataengine / dude_tts /runtime.py
penguinfish1688's picture
Support user-provided speaker audio with a bundled reference encoder
ac37044 verified
Raw History Blame Contribute Delete
7.64 kB
"""XML + two reference voices -> normalized 24 kHz stereo audio."""
from contextlib import nullcontext
from dataclasses import dataclass
import json
import math
from pathlib import Path
from types import SimpleNamespace
import numpy as np
import soundfile as sf
import torch
from accelerate import init_empty_weights
from huggingface_hub import snapshot_download
from safetensors.torch import load_file
from transformers import AutoTokenizer
from .audio import normalize_for_listening
from .frontend import duplex_prefix
from .interleave import EVENTS
from .model import DualChannelTTS
from .voice import ReferenceVoiceEncoder
@dataclass
class Generation:
audio: np.ndarray
sample_rate: int
eos: tuple[bool, bool]
frames: tuple[int, int]
normalization: list[dict]
def save(self, path):
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
sf.write(path, self.audio, self.sample_rate, subtype='PCM_16')
return path
class DuDE:
"""Load only the public release; no training cache or research checkout."""
@classmethod
def from_pretrained(cls, model_id='penguinfish1688/duplexdataengine', *, device='cuda', revision=None):
root = Path(model_id)
if not root.is_dir():
root = Path(snapshot_download(model_id, revision=revision,
allow_patterns=['*.json', '*.txt', '*.safetensors', 'speech_tokenizer/*',
'speaker_encoder/*', 'voices/*']))
return cls(root, device=device)
def __init__(self, root, *, device='cuda'):
from qwen_tts.core.models.configuration_qwen3_tts import Qwen3TTSConfig
from qwen_tts.core.models.modeling_qwen3_tts import Qwen3TTSTalkerForConditionalGeneration
from qwen_tts.inference.qwen3_tts_tokenizer import Qwen3TTSTokenizer
self.root = Path(root)
self.device = torch.device(device)
if self.device.type == 'cuda' and not torch.cuda.is_available():
raise RuntimeError('CUDA is unavailable. Install a CUDA PyTorch build or select device="cpu" (slow).')
self.release = json.loads((self.root/'release.json').read_text())
self.events = json.loads((self.root/'event-token-map.json').read_text())
self.voices = json.loads((self.root/'voices/presets.json').read_text())
self._voice_encoder = None
self._voice_cache = {}
config = Qwen3TTSConfig.from_pretrained(self.root, local_files_only=True)
config.talker_config._attn_implementation = 'sdpa'
config.talker_config.code_predictor_config._attn_implementation = 'sdpa'
self.tokenizer = AutoTokenizer.from_pretrained(self.root, local_files_only=True)
vocab = set(self.tokenizer.get_vocab().values())
if (set(self.events) != set(EVENTS) or len(set(self.events.values())) != len(EVENTS)
or vocab & set(self.events.values())
or any(type(i) is not int or not 0 <= i < config.talker_config.text_vocab_size
for i in self.events.values())):
raise ValueError('Invalid release event-token map')
# Keep nonpersistent rotary buffers real while allocating parameters on
# meta; safetensors supplies every actual parameter without a random model copy.
with init_empty_weights():
talker = Qwen3TTSTalkerForConditionalGeneration(config.talker_config)
model = DualChannelTTS(SimpleNamespace(model=SimpleNamespace(talker=talker, config=config)))
state = load_file(str(self.root/'model.safetensors'))
model.load_state_dict(state, strict=True, assign=True)
self.model = model.to(self.device).eval()
self.codec = Qwen3TTSTokenizer.from_pretrained(self.root/'speech_tokenizer',
device_map=str(self.device), dtype=torch.bfloat16 if self.device.type == 'cuda' else torch.float32,
attn_implementation='sdpa', local_files_only=True)
def _text(self, text):
return self.tokenizer.encode(text, add_special_tokens=False)
def _event(self, event):
return [self.events[event]]
def voice_embedding(self, reference):
"""Accept an included voice name or a path to a single-speaker recording."""
if isinstance(reference, str) and reference in self.voices:
return self.voices[reference]['embedding']
if not isinstance(reference, (str, Path)):
raise ValueError('A voice must be an audio file path or an included voice name.')
path = Path(reference).expanduser().resolve()
if not path.is_file():
raise ValueError(f'Reference audio {str(reference)!r} does not exist. '
f'Provide an audio file or choose from {list(self.voices)}.')
stat = path.stat()
key = (str(path), stat.st_mtime_ns, stat.st_size)
if key not in self._voice_cache:
if self._voice_encoder is None:
self._voice_encoder = ReferenceVoiceEncoder(self.root/'speaker_encoder', self.device)
self._voice_cache[key] = self._voice_encoder.encode(path)
return self._voice_cache[key]
def conditioning(self, xml, voice_a, voice_b):
if not isinstance(xml, str) or not xml.strip() or len(xml) > 30000:
raise ValueError('Enter a nonempty dialogue XML, up to 30,000 characters.')
header = dict(text_ids=self._text('<|im_start|>assistant\n<|im_end|>\n<|im_start|>assistant\n'),
instruct_ids=[], language='English', speaker='Ryan')
channels = {}
for channel, voice in zip('AB', [voice_a, voice_b]):
channels[channel] = dict(header, voice_embedding=self.voice_embedding(voice))
record = dict(xml=xml, channels=channels)
conditioning = duplex_prefix(record, self.model.config, self._text, self._event)
if len(conditioning['text_ids']) > 12000:
raise ValueError('Dialogue is too long for this inference interface; use fewer turns.')
return conditioning
@torch.inference_mode()
def generate(self, xml, voice_a, voice_b, *, seed=20260922, max_seconds=150.,
temperature=.9, top_k=50, greedy=False):
if not math.isfinite(max_seconds) or not 1 <= max_seconds <= 300:
raise ValueError('The safety limit must be between 1 and 300 seconds.')
conditioning = self.conditioning(xml, voice_a, voice_b)
devices = [self.device.index or 0] if self.device.type == 'cuda' else []
autocast = torch.autocast('cuda', dtype=torch.bfloat16) if devices else nullcontext()
with torch.random.fork_rng(devices=devices), autocast:
torch.manual_seed(int(seed))
arrays, ended = self.model.generate_pair([], _conditioning=conditioning,
max_frames=math.ceil(max_seconds*12.5), do_sample=not greedy,
temperature=float(temperature), top_k=int(top_k),
subtalker_temperature=float(temperature), subtalker_top_k=int(top_k))
if any(a is None for a in arrays):
raise RuntimeError('The model produced an empty audio lane; try a different seed.')
waves, rate = self.codec.decode([{'audio_codes': a} for a in arrays])
normalized, reports = zip(*(normalize_for_listening(w, rate) for w in waves))
stereo = np.zeros((max(map(len, normalized)), 2), dtype=np.float32)
for lane, wave in enumerate(normalized):
stereo[:len(wave), lane] = wave
return Generation(stereo, rate, tuple(bool(v) for v in ended),
tuple(len(a) for a in arrays), list(reports))