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