File size: 7,643 Bytes
ac37044
4123b95
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
ac37044
4123b95
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
ac37044
 
4123b95
 
 
 
 
 
 
 
 
 
 
 
 
ac37044
 
4123b95
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
ac37044
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4123b95
 
 
 
 
 
 
ac37044
4123b95
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
"""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))