Prompt48 commited on
Commit
bfbecd0
·
verified ·
1 Parent(s): f5b2a58

Upload edit\Qwen3-TTS-test\qwen_tts\core\tokenizer_25hz\modeling_qwen3_tts_tokenizer_v1.py with huggingface_hub

Browse files
edit//Qwen3-TTS-test//qwen_tts//core//tokenizer_25hz//modeling_qwen3_tts_tokenizer_v1.py ADDED
@@ -0,0 +1,1529 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # coding=utf-8
2
+ # Copyright 2026 The Qwen team, Alibaba Group and the HuggingFace Inc. team. All rights reserved.
3
+ #
4
+ # Licensed under the Apache License, Version 2.0 (the "License");
5
+ # you may not use this file except in compliance with the License.
6
+ # You may obtain a copy of the License at
7
+ #
8
+ # http://www.apache.org/licenses/LICENSE-2.0
9
+ #
10
+ # Unless required by applicable law or agreed to in writing, software
11
+ # distributed under the License is distributed on an "AS IS" BASIS,
12
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ # See the License for the specific language governing permissions and
14
+ # limitations under the License.
15
+ """PyTorch Qwen3TTSTokenizerV1 model."""
16
+
17
+ import math
18
+ from dataclasses import dataclass
19
+ from typing import Optional, Union, List
20
+
21
+ import numpy as np
22
+ import torch
23
+ from torch import nn
24
+ from torch.nn import Parameter
25
+ from torch.nn import functional as F
26
+ from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel
27
+ from transformers.utils import ModelOutput, auto_docstring, logging
28
+ from transformers.utils.hub import cached_file
29
+
30
+ from torch.nn.utils.rnn import pad_sequence
31
+
32
+ from .vq.whisper_encoder import get_mel_audio, get_T_after_cnn
33
+ from .vq.speech_vq import WhisperEncoderVQ, XVectorExtractor
34
+
35
+ from .configuration_qwen3_tts_tokenizer_v1 import (
36
+ Qwen3TTSTokenizerV1Config,
37
+ Qwen3TTSTokenizerV1EncoderConfig,
38
+ Qwen3TTSTokenizerV1DecoderConfig,
39
+ Qwen3TTSTokenizerV1DecoderBigVGANConfig,
40
+ Qwen3TTSTokenizerV1DecoderDiTConfig
41
+ )
42
+
43
+ logger = logging.get_logger(__name__)
44
+
45
+
46
+ @dataclass
47
+ @auto_docstring
48
+ class Qwen3TTSTokenizerV1EncoderOutput(ModelOutput):
49
+ r"""
50
+ audio_codes (`List[torch.LongTensor]`):
51
+ Discret code embeddings computed using `model.encode`, each tensor has shape (codes_length_i,).
52
+ xvectors (`List[torch.FloatTensor]`):
53
+ X-vector embeddings computed using `model.encode`, each tensor has shape (xvector_dim,).
54
+ ref_mels (`List[torch.FloatTensor]`):
55
+ Reference mel spectrogram computed using `model.encode`, each tensor has shape (mel_length_i, mel_dim,).
56
+ """
57
+
58
+ audio_codes: List[torch.LongTensor] = None
59
+ xvectors: List[torch.FloatTensor] = None
60
+ ref_mels: List[torch.FloatTensor] = None
61
+
62
+
63
+ @dataclass
64
+ @auto_docstring
65
+ class Qwen3TTSTokenizerV1DecoderOutput(ModelOutput):
66
+ r"""
67
+ audio_values (`List[torch.FloatTensor]`):
68
+ Decoded audio values, obtained using the decoder part of Qwen3TTSTokenizerV1.
69
+ Each tensor has shape (segment_length_i).
70
+ """
71
+
72
+ audio_values: List[torch.FloatTensor] = None
73
+
74
+
75
+ @auto_docstring
76
+ class Qwen3TTSTokenizerV1DecoderPreTrainedModel(PreTrainedModel):
77
+ config: Qwen3TTSTokenizerV1DecoderConfig
78
+ base_model_prefix = "model"
79
+ supports_gradient_checkpointing = True
80
+ _skip_keys_device_placement = "past_key_values"
81
+ _supports_flash_attn = True
82
+ _supports_sdpa = True
83
+ _can_compile_fullgraph = False
84
+ _supports_attention_backend = True
85
+
86
+
87
+ @auto_docstring
88
+ class Qwen3TTSTokenizerV1EncoderPreTrainedModel(PreTrainedModel):
89
+ config: Qwen3TTSTokenizerV1EncoderConfig
90
+ base_model_prefix = "model"
91
+ supports_gradient_checkpointing = True
92
+ _skip_keys_device_placement = "past_key_values"
93
+ _supports_flash_attn = True
94
+ _supports_sdpa = True
95
+ _can_compile_fullgraph = False
96
+ _supports_attention_backend = True
97
+
98
+
99
+ class Qwen3TTSTokenizerV1DecoderDiTRotaryEmbedding(nn.Module):
100
+ inv_freq: torch.Tensor # fix linting for `register_buffer`
101
+
102
+ def __init__(self, dim, base=10000):
103
+ super().__init__()
104
+
105
+ inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2).float() / dim))
106
+ self.register_buffer("inv_freq", inv_freq)
107
+
108
+ def forward(self, x):
109
+ batch_size, seq_len = x.shape[0], x.shape[1]
110
+ t = torch.arange(seq_len, device=x.device)
111
+ device_type = x.device.type
112
+ device_type = device_type if device_type != "mps" else "cpu"
113
+ with torch.autocast(device_type=device_type, enabled=False):
114
+ freqs = t.unsqueeze(1).float() @ self.inv_freq.unsqueeze(0).float()
115
+ freqs = torch.stack((freqs, freqs), dim=-1)
116
+ freqs = freqs.reshape(*freqs.shape[:-2], -1)
117
+ freqs = freqs.repeat(batch_size, *([1] * freqs.dim()))
118
+ cos = freqs.cos()
119
+ sin = freqs.sin()
120
+
121
+ return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype)
122
+
123
+
124
+ class TimeDelayNetBlock(nn.Module):
125
+ def __init__(
126
+ self,
127
+ in_channels,
128
+ out_channels,
129
+ kernel_size,
130
+ dilation,
131
+ ):
132
+ super().__init__()
133
+ self.conv = nn.Conv1d(
134
+ in_channels=in_channels,
135
+ out_channels=out_channels,
136
+ kernel_size=kernel_size,
137
+ dilation=dilation,
138
+ padding="same",
139
+ padding_mode="reflect",
140
+ )
141
+ self.activation = nn.ReLU()
142
+
143
+ def forward(self, hidden_states: torch.Tensor):
144
+ return self.activation(self.conv(hidden_states))
145
+
146
+
147
+ class Res2NetBlock(torch.nn.Module):
148
+ def __init__(self, in_channels, out_channels, scale=8, kernel_size=3, dilation=1):
149
+ super().__init__()
150
+
151
+ in_channel = in_channels // scale
152
+ hidden_channel = out_channels // scale
153
+
154
+ self.blocks = nn.ModuleList(
155
+ [
156
+ TimeDelayNetBlock(
157
+ in_channel,
158
+ hidden_channel,
159
+ kernel_size=kernel_size,
160
+ dilation=dilation,
161
+ )
162
+ for i in range(scale - 1)
163
+ ]
164
+ )
165
+ self.scale = scale
166
+
167
+ def forward(self, hidden_states):
168
+ outputs = []
169
+ for i, hidden_part in enumerate(torch.chunk(hidden_states, self.scale, dim=1)):
170
+ if i == 0:
171
+ output_part = hidden_part
172
+ elif i == 1:
173
+ output_part = self.blocks[i - 1](hidden_part)
174
+ else:
175
+ output_part = self.blocks[i - 1](hidden_part + output_part)
176
+ outputs.append(output_part)
177
+ output = torch.cat(outputs, dim=1)
178
+ return output
179
+
180
+
181
+ class SqueezeExcitationBlock(nn.Module):
182
+ def __init__(self, in_channels, se_channels, out_channels):
183
+ super().__init__()
184
+
185
+ self.conv1 = nn.Conv1d(
186
+ in_channels=in_channels,
187
+ out_channels=se_channels,
188
+ kernel_size=1,
189
+ padding="same",
190
+ padding_mode="reflect",
191
+ )
192
+ self.relu = nn.ReLU(inplace=True)
193
+ self.conv2 = nn.Conv1d(
194
+ in_channels=se_channels,
195
+ out_channels=out_channels,
196
+ kernel_size=1,
197
+ padding="same",
198
+ padding_mode="reflect",
199
+ )
200
+ self.sigmoid = nn.Sigmoid()
201
+
202
+ def forward(self, hidden_states):
203
+ hidden_states_mean = hidden_states.mean(dim=2, keepdim=True)
204
+
205
+ hidden_states_mean = self.relu(self.conv1(hidden_states_mean))
206
+ hidden_states_mean = self.sigmoid(self.conv2(hidden_states_mean))
207
+
208
+ return hidden_states * hidden_states_mean
209
+
210
+
211
+ class AttentiveStatisticsPooling(nn.Module):
212
+ """This class implements an attentive statistic pooling layer for each channel.
213
+ It returns the concatenated mean and std of the input tensor.
214
+ """
215
+
216
+ def __init__(self, channels, attention_channels=128):
217
+ super().__init__()
218
+
219
+ self.eps = 1e-12
220
+ self.tdnn = TimeDelayNetBlock(channels * 3, attention_channels, 1, 1)
221
+ self.tanh = nn.Tanh()
222
+ self.conv = nn.Conv1d(
223
+ in_channels=attention_channels,
224
+ out_channels=channels,
225
+ kernel_size=1,
226
+ padding="same",
227
+ padding_mode="reflect",
228
+ )
229
+
230
+ def _length_to_mask(self, length, max_len=None, dtype=None, device=None):
231
+ """Creates a binary mask for each sequence.
232
+
233
+ Reference: https://discuss.pytorch.org/t/how-to-generate-variable-length-mask/23397/3
234
+
235
+ Arguments
236
+ ---------
237
+ length : torch.LongTensor
238
+ Containing the length of each sequence in the batch. Must be 1D.
239
+ max_len : int
240
+ Max length for the mask, also the size of the second dimension.
241
+ dtype : torch.dtype, default: None
242
+ The dtype of the generated mask.
243
+ device: torch.device, default: None
244
+ The device to put the mask variable.
245
+
246
+ Returns
247
+ -------
248
+ mask : tensor
249
+ The binary mask.
250
+ """
251
+
252
+ if max_len is None:
253
+ max_len = length.max().long().item() # using arange to generate mask
254
+ mask = torch.arange(max_len, device=length.device, dtype=length.dtype).expand(
255
+ len(length), max_len
256
+ ) < length.unsqueeze(1)
257
+
258
+ mask = torch.as_tensor(mask, dtype=dtype, device=device)
259
+ return mask
260
+
261
+ def _compute_statistics(self, x, m, dim=2):
262
+ mean = (m * x).sum(dim)
263
+ std = torch.sqrt((m * (x - mean.unsqueeze(dim)).pow(2)).sum(dim).clamp(self.eps))
264
+ return mean, std
265
+
266
+ def forward(self, hidden_states):
267
+ seq_length = hidden_states.shape[-1]
268
+ lengths = torch.ones(hidden_states.shape[0], device=hidden_states.device)
269
+
270
+ # Make binary mask of shape [N, 1, L]
271
+ mask = self._length_to_mask(
272
+ lengths * seq_length, max_len=seq_length, dtype=hidden_states.dtype, device=hidden_states.device
273
+ )
274
+ mask = mask.unsqueeze(1)
275
+
276
+ # Expand the temporal context of the pooling layer by allowing the
277
+ # self-attention to look at global properties of the utterance.
278
+ total = mask.sum(dim=2, keepdim=True)
279
+
280
+ mean, std = self._compute_statistics(hidden_states, mask / total)
281
+ mean = mean.unsqueeze(2).repeat(1, 1, seq_length)
282
+ std = std.unsqueeze(2).repeat(1, 1, seq_length)
283
+ attention = torch.cat([hidden_states, mean, std], dim=1)
284
+
285
+ # Apply layers
286
+ attention = self.conv(self.tanh(self.tdnn(attention)))
287
+
288
+ # Filter out zero-paddings
289
+ attention = attention.masked_fill(mask == 0, float("-inf"))
290
+
291
+ attention = F.softmax(attention, dim=2)
292
+ mean, std = self._compute_statistics(hidden_states, attention)
293
+ # Append mean and std of the batch
294
+ pooled_stats = torch.cat((mean, std), dim=1)
295
+ pooled_stats = pooled_stats.unsqueeze(2)
296
+
297
+ return pooled_stats
298
+
299
+
300
+ class SqueezeExcitationRes2NetBlock(nn.Module):
301
+ """An implementation of building block in ECAPA-TDNN, i.e.,
302
+ TDNN-Res2Net-TDNN-SqueezeExcitationBlock.
303
+ """
304
+
305
+ def __init__(
306
+ self,
307
+ in_channels,
308
+ out_channels,
309
+ res2net_scale=8,
310
+ se_channels=128,
311
+ kernel_size=1,
312
+ dilation=1,
313
+ ):
314
+ super().__init__()
315
+ self.out_channels = out_channels
316
+ self.tdnn1 = TimeDelayNetBlock(
317
+ in_channels,
318
+ out_channels,
319
+ kernel_size=1,
320
+ dilation=1,
321
+ )
322
+ self.res2net_block = Res2NetBlock(out_channels, out_channels, res2net_scale, kernel_size, dilation)
323
+ self.tdnn2 = TimeDelayNetBlock(
324
+ out_channels,
325
+ out_channels,
326
+ kernel_size=1,
327
+ dilation=1,
328
+ )
329
+ self.se_block = SqueezeExcitationBlock(out_channels, se_channels, out_channels)
330
+
331
+ def forward(self, hidden_state):
332
+ residual = hidden_state
333
+
334
+ hidden_state = self.tdnn1(hidden_state)
335
+ hidden_state = self.res2net_block(hidden_state)
336
+ hidden_state = self.tdnn2(hidden_state)
337
+ hidden_state = self.se_block(hidden_state)
338
+
339
+ return hidden_state + residual
340
+
341
+
342
+ class ECAPA_TimeDelayNet(torch.nn.Module):
343
+ """An implementation of the speaker embedding model in a paper.
344
+ "ECAPA-TDNN: Emphasized Channel Attention, Propagation and Aggregation in
345
+ TDNN Based Speaker Verification" (https://huggingface.co/papers/2005.07143).
346
+ """
347
+
348
+ def __init__(self, config: Qwen3TTSTokenizerV1DecoderBigVGANConfig):
349
+ super().__init__()
350
+ if len(config.enc_channels) != len(config.enc_kernel_sizes) or len(config.enc_channels) != len(
351
+ config.enc_dilations
352
+ ):
353
+ raise ValueError("enc_channels, enc_kernel_sizes and enc_dilations should have same length")
354
+ self.channels = config.enc_channels
355
+ self.blocks = nn.ModuleList()
356
+
357
+ # The initial TDNN layer
358
+ self.blocks.append(
359
+ TimeDelayNetBlock(
360
+ config.mel_dim,
361
+ config.enc_channels[0],
362
+ config.enc_kernel_sizes[0],
363
+ config.enc_dilations[0],
364
+ )
365
+ )
366
+
367
+ # SE-Res2Net layers
368
+ for i in range(1, len(config.enc_channels) - 1):
369
+ self.blocks.append(
370
+ SqueezeExcitationRes2NetBlock(
371
+ config.enc_channels[i - 1],
372
+ config.enc_channels[i],
373
+ res2net_scale=config.enc_res2net_scale,
374
+ se_channels=config.enc_se_channels,
375
+ kernel_size=config.enc_kernel_sizes[i],
376
+ dilation=config.enc_dilations[i],
377
+ )
378
+ )
379
+
380
+ # Multi-layer feature aggregation
381
+ self.mfa = TimeDelayNetBlock(
382
+ config.enc_channels[-1],
383
+ config.enc_channels[-1],
384
+ config.enc_kernel_sizes[-1],
385
+ config.enc_dilations[-1],
386
+ )
387
+
388
+ # Attentive Statistical Pooling
389
+ self.asp = AttentiveStatisticsPooling(
390
+ config.enc_channels[-1],
391
+ attention_channels=config.enc_attention_channels,
392
+ )
393
+
394
+ # Final linear transformation
395
+ self.fc = nn.Conv1d(
396
+ in_channels=config.enc_channels[-1] * 2,
397
+ out_channels=config.enc_dim,
398
+ kernel_size=1,
399
+ padding="same",
400
+ padding_mode="reflect",
401
+ )
402
+
403
+ def forward(self, hidden_states):
404
+ # Minimize transpose for efficiency
405
+ hidden_states = hidden_states.transpose(1, 2)
406
+
407
+ hidden_states_list = []
408
+ for layer in self.blocks:
409
+ hidden_states = layer(hidden_states)
410
+ hidden_states_list.append(hidden_states)
411
+
412
+ # Multi-layer feature aggregation
413
+ hidden_states = torch.cat(hidden_states_list[1:], dim=1)
414
+ hidden_states = self.mfa(hidden_states)
415
+
416
+ # Attentive Statistical Pooling
417
+ hidden_states = self.asp(hidden_states)
418
+
419
+ # Final linear transformation
420
+ hidden_states = self.fc(hidden_states)
421
+
422
+ hidden_states = hidden_states.squeeze(-1)
423
+ return hidden_states
424
+
425
+
426
+ class DiTInputEmbedding(nn.Module):
427
+ def __init__(self, config: Qwen3TTSTokenizerV1DecoderBigVGANConfig):
428
+ super().__init__()
429
+ self.proj = nn.Linear(
430
+ config.mel_dim + config.enc_dim + config.enc_emb_dim + config.emb_dim,
431
+ config.hidden_size,
432
+ )
433
+ self.spk_encoder = ECAPA_TimeDelayNet(config)
434
+
435
+ def forward(
436
+ self,
437
+ hidden_states: torch.Tensor,
438
+ speaker_embedding: torch.Tensor,
439
+ condition_vector: torch.Tensor,
440
+ code_embed: torch.Tensor,
441
+ drop_audio_cond: Optional[bool] = False,
442
+ code_embed_uncond: Optional[bool] = None,
443
+ apply_cfg: Optional[bool] = True,
444
+ ):
445
+ if apply_cfg:
446
+ hidden_states = torch.cat([hidden_states, hidden_states], dim=0)
447
+ speaker_embedding = torch.cat([speaker_embedding, torch.zeros_like(speaker_embedding)], dim=0)
448
+ condition_vector = torch.cat([condition_vector, torch.zeros_like(condition_vector)], dim=0)
449
+ code_embed = torch.cat([code_embed, code_embed_uncond], dim=0)
450
+ elif drop_audio_cond: # cfg for cond audio
451
+ condition_vector = torch.zeros_like(condition_vector)
452
+ speaker_embedding = torch.zeros_like(speaker_embedding)
453
+ condition_vector = self.spk_encoder(condition_vector).unsqueeze(1).repeat(1, hidden_states.size(1), 1)
454
+ hidden_states = self.proj(torch.cat((hidden_states, condition_vector, code_embed, speaker_embedding), dim=-1))
455
+
456
+ return hidden_states
457
+
458
+
459
+ # Transformer backbone using DiT blocks
460
+ class DiTCodecEmbedding(nn.Module):
461
+ def __init__(self, codec_num_embeds, codec_dim, repeats):
462
+ super().__init__()
463
+ self.repeats = repeats
464
+ self.codec_embed = nn.Embedding(codec_num_embeds + 1, codec_dim)
465
+
466
+ def forward(self, code, drop_code=False):
467
+ if drop_code:
468
+ code = torch.zeros_like(code)
469
+ code_embed = self.codec_embed(code)
470
+
471
+ code_embed = torch.repeat_interleave(code_embed, repeats=self.repeats, dim=1)
472
+ return code_embed
473
+
474
+
475
+ # AdaLayerNormZero
476
+ # return with modulated x for attn input, and params for later mlp modulation
477
+ class AdaLayerNormZero(nn.Module):
478
+ def __init__(self, dim):
479
+ super().__init__()
480
+
481
+ self.silu = nn.SiLU()
482
+ self.linear = nn.Linear(dim, dim * 6)
483
+
484
+ self.norm = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
485
+
486
+ def forward(self, hidden_states, emb=None):
487
+ emb = self.linear(self.silu(emb))
488
+ shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = torch.chunk(emb, 6, dim=1)
489
+
490
+ hidden_states = self.norm(hidden_states) * (1 + scale_msa[:, None]) + shift_msa[:, None]
491
+ return hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp
492
+
493
+
494
+ # AdaLayerNormZero for final layer
495
+ # return only with modulated x for attn input, cuz no more mlp modulation
496
+ class AdaLayerNormZero_Final(nn.Module):
497
+ def __init__(self, dim):
498
+ super().__init__()
499
+
500
+ self.silu = nn.SiLU()
501
+ self.linear = nn.Linear(dim, dim * 2)
502
+
503
+ self.norm = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
504
+
505
+ def forward(self, hidden_states, emb):
506
+ emb = self.linear(self.silu(emb))
507
+ scale, shift = torch.chunk(emb, 2, dim=1)
508
+
509
+ hidden_states = self.norm(hidden_states) * (1 + scale)[:, None, :] + shift[:, None, :]
510
+ return hidden_states
511
+
512
+
513
+ # FeedForward
514
+ class DiTMLP(nn.Module):
515
+ def __init__(self, dim, mult=4, dropout=0.0):
516
+ super().__init__()
517
+ inner_dim = int(dim * mult)
518
+
519
+ self.ff = nn.ModuleList(
520
+ [
521
+ nn.Linear(dim, inner_dim),
522
+ nn.GELU(approximate="tanh"),
523
+ nn.Dropout(dropout),
524
+ nn.Linear(inner_dim, dim),
525
+ ]
526
+ )
527
+
528
+ def forward(self, hidden_states):
529
+ for layer in self.ff:
530
+ hidden_states = layer(hidden_states)
531
+ return hidden_states
532
+
533
+
534
+ # Modified from Llama with a different rotate function, will fixed in next release
535
+ def apply_rotary_pos_emb(q, k, cos, sin, position_ids=None, unsqueeze_dim=1):
536
+ """Applies Rotary Position Embedding to the query and key tensors.
537
+
538
+ Args:
539
+ q (`torch.Tensor`): The query tensor.
540
+ k (`torch.Tensor`): The key tensor.
541
+ cos (`torch.Tensor`): The cosine part of the rotary embedding.
542
+ sin (`torch.Tensor`): The sine part of the rotary embedding.
543
+ position_ids (`torch.Tensor`, *optional*):
544
+ Deprecated and unused.
545
+ unsqueeze_dim (`int`, *optional*, defaults to 1):
546
+ The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and
547
+ sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note
548
+ that cos[position_ids] and sin[position_ids] have the shape [batch_size, seq_len, head_dim]. Then, if q and
549
+ k have the shape [batch_size, heads, seq_len, head_dim], then setting unsqueeze_dim=1 makes
550
+ cos[position_ids] and sin[position_ids] broadcastable to the shapes of q and k. Similarly, if q and k have
551
+ the shape [batch_size, seq_len, heads, head_dim], then set unsqueeze_dim=2.
552
+ Returns:
553
+ `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding.
554
+ """
555
+
556
+ def rotate_half_codec(x):
557
+ # x = rearrange(x, "... (d r) -> ... d r", r=2)
558
+ x = x.reshape(*x.shape[:-1], -1, 2)
559
+ x1, x2 = x.unbind(dim=-1)
560
+ x = torch.stack((-x2, x1), dim=-1)
561
+ return x.reshape(*x.shape[:-2], -1)
562
+
563
+ cos = cos.unsqueeze(unsqueeze_dim)
564
+ sin = sin.unsqueeze(unsqueeze_dim)
565
+ q_embed = (q * cos) + (rotate_half_codec(q) * sin)
566
+ k_embed = (k * cos) + (rotate_half_codec(k) * sin)
567
+ return q_embed, k_embed
568
+
569
+
570
+ class DiTAttention(nn.Module):
571
+ def __init__(self, config: Qwen3TTSTokenizerV1DecoderBigVGANConfig):
572
+ super().__init__()
573
+
574
+ self.config = config
575
+ self.dim = config.hidden_size
576
+ self.heads = config.num_attention_heads
577
+ self.inner_dim = config.head_dim * config.num_attention_heads
578
+ self.dropout = config.dropout
579
+ self.is_causal = False
580
+
581
+ self.to_q = nn.Linear(config.hidden_size, self.inner_dim)
582
+ self.to_k = nn.Linear(config.hidden_size, self.inner_dim)
583
+ self.to_v = nn.Linear(config.hidden_size, self.inner_dim)
584
+
585
+ self.to_out = nn.ModuleList([nn.Linear(self.inner_dim, config.hidden_size), nn.Dropout(config.dropout)])
586
+
587
+ def forward(
588
+ self,
589
+ hidden_states, # noised input x
590
+ position_embeddings=None, # rotary position embedding for x
591
+ attention_mask=None,
592
+ ) -> torch.Tensor:
593
+ batch_size = hidden_states.shape[0]
594
+
595
+ # `sample` projections.
596
+ query = self.to_q(hidden_states)
597
+ key = self.to_k(hidden_states)
598
+ value = self.to_v(hidden_states)
599
+
600
+ # attention
601
+ inner_dim = key.shape[-1]
602
+ head_dim = inner_dim // self.heads
603
+ query = query.view(batch_size, -1, self.heads, head_dim).transpose(1, 2)
604
+ key = key.view(batch_size, -1, self.heads, head_dim).transpose(1, 2)
605
+ value = value.view(batch_size, -1, self.heads, head_dim).transpose(1, 2)
606
+
607
+ # apply rotary position embedding
608
+ # Due to training process, only first head is applied with RoPE, will be fixed at next release
609
+ cos, sin = position_embeddings
610
+ query, key = apply_rotary_pos_emb(query, key, cos, sin)
611
+
612
+ attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation]
613
+ attention_weights, _ = attention_interface(
614
+ self,
615
+ query,
616
+ key,
617
+ value,
618
+ attention_mask=attention_mask,
619
+ is_causal=False,
620
+ )
621
+
622
+ # mask. e.g. inference got a batch with different target durations, mask out the padding
623
+ attention_weights = attention_weights.reshape(batch_size, -1, self.heads * head_dim)
624
+ attention_weights = attention_weights.to(query.dtype)
625
+
626
+ # linear proj
627
+ attention_output = self.to_out[0](attention_weights)
628
+ attention_output = self.to_out[1](attention_output)
629
+
630
+ return attention_output
631
+
632
+
633
+ # time step conditioning embedding
634
+ class SinusPositionEmbedding(nn.Module):
635
+ def __init__(self, dim):
636
+ super().__init__()
637
+ self.dim = dim
638
+
639
+ def forward(self, hidden_states, scale=1000):
640
+ device = hidden_states.device
641
+ half_dim = self.dim // 2
642
+ emb = math.log(10000) / (half_dim - 1)
643
+ emb = torch.exp(torch.arange(half_dim, device=device).float() * -emb)
644
+ emb = scale * hidden_states.unsqueeze(1) * emb.unsqueeze(0)
645
+ emb = torch.cat((emb.sin(), emb.cos()), dim=-1)
646
+ return emb.type_as(hidden_states)
647
+
648
+
649
+ class DiTTimestepEmbedding(nn.Module):
650
+ def __init__(self, dim, freq_embed_dim=256):
651
+ super().__init__()
652
+ self.time_embed = SinusPositionEmbedding(freq_embed_dim)
653
+ self.time_mlp = nn.ModuleList([nn.Linear(freq_embed_dim, dim), nn.SiLU(), nn.Linear(dim, dim)])
654
+
655
+ def forward(self, timestep):
656
+ time_hidden = self.time_embed(timestep)
657
+ time_hidden = time_hidden.to(timestep.dtype)
658
+ for layer in self.time_mlp:
659
+ time_hidden = layer(time_hidden) # b d
660
+ return time_hidden
661
+
662
+
663
+ class DiTDecoderLayer(nn.Module):
664
+ def __init__(self, config: Qwen3TTSTokenizerV1DecoderBigVGANConfig, look_ahead_block=0, look_backward_block=0):
665
+ super().__init__()
666
+ self.attn_norm = AdaLayerNormZero(config.hidden_size)
667
+
668
+ self.attn = DiTAttention(config)
669
+ self.look_ahead_block = look_ahead_block
670
+ self.look_backward_block = look_backward_block
671
+ self.ff_norm = nn.LayerNorm(config.hidden_size, elementwise_affine=False, eps=1e-6)
672
+ self.ff = DiTMLP(dim=config.hidden_size, mult=config.ff_mult, dropout=config.dropout)
673
+
674
+ def forward(
675
+ self, hidden_states, timestep, position_embeddings=None, block_diff=None
676
+ ): # x: noised input, t: time embedding
677
+ # pre-norm & modulation for attention input
678
+ norm, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.attn_norm(hidden_states, emb=timestep)
679
+
680
+ # attention
681
+ attn_output = self.attn(
682
+ hidden_states=norm,
683
+ position_embeddings=position_embeddings,
684
+ attention_mask=(block_diff >= -float(self.look_backward_block))
685
+ & (block_diff <= float(self.look_ahead_block)),
686
+ )
687
+
688
+ # process attention output for input x
689
+ hidden_states = hidden_states + gate_msa.unsqueeze(1) * attn_output
690
+
691
+ norm = self.ff_norm(hidden_states) * (1 + scale_mlp[:, None]) + shift_mlp[:, None]
692
+ ff_output = self.ff(norm)
693
+ hidden_states = hidden_states + gate_mlp.unsqueeze(1) * ff_output
694
+
695
+ return hidden_states
696
+
697
+
698
+ class SnakeBeta(nn.Module):
699
+ """
700
+ A modified Snake function which uses separate parameters for the magnitude of the periodic components
701
+ Shape:
702
+ - Input: (B, C, T)
703
+ - Output: (B, C, T), same shape as the input
704
+ Parameters:
705
+ - alpha - trainable parameter that controls frequency
706
+ - beta - trainable parameter that controls magnitude
707
+ References:
708
+ - This activation function is a modified version based on this paper by Liu Ziyin, Tilman Hartwig, Masahito Ueda:
709
+ https://huggingface.co/papers/2006.08195
710
+ """
711
+
712
+ def __init__(self, in_features, alpha=1.0):
713
+ super().__init__()
714
+ self.in_features = in_features
715
+
716
+ # initialize alpha
717
+ self.alpha = Parameter(torch.zeros(in_features) * alpha)
718
+ self.beta = Parameter(torch.zeros(in_features) * alpha)
719
+
720
+ self.no_div_by_zero = 0.000000001
721
+
722
+ def forward(self, hidden_states):
723
+ """
724
+ Forward pass of the function.
725
+ Applies the function to the input elementwise.
726
+ SnakeBeta ∶= x + 1/b * sin^2 (xa)
727
+ """
728
+ alpha = self.alpha.unsqueeze(0).unsqueeze(-1) # line up with x to [B, C, T]
729
+ beta = self.beta.unsqueeze(0).unsqueeze(-1)
730
+ alpha = torch.exp(alpha)
731
+ beta = torch.exp(beta)
732
+ hidden_states = hidden_states + (1.0 / (beta + self.no_div_by_zero)) * torch.pow(
733
+ torch.sin(hidden_states * alpha), 2
734
+ )
735
+
736
+ return hidden_states
737
+
738
+
739
+ def kaiser_sinc_filter1d(cutoff, half_width, kernel_size):
740
+ """Generates a 1D Kaiser-windowed sinc filter.
741
+
742
+ Args:
743
+ cutoff (float): Normalized cutoff frequency (0 to 0.5).
744
+ half_width (float): Transition bandwidth.
745
+ kernel_size (int): Number of filter taps.
746
+
747
+ Returns:
748
+ torch.Tensor: A tensor of shape (1, 1, kernel_size) representing the filter.
749
+ """
750
+ is_even = kernel_size % 2 == 0
751
+ half_size = kernel_size // 2
752
+
753
+ # Compute Kaiser window parameters
754
+ delta_f = 4 * half_width
755
+ attenuation = 2.285 * (half_size - 1) * math.pi * delta_f + 7.95
756
+
757
+ if attenuation > 50.0:
758
+ beta = 0.1102 * (attenuation - 8.7)
759
+ elif attenuation >= 21.0:
760
+ beta = 0.5842 * (attenuation - 21) ** 0.4 + 0.07886 * (attenuation - 21.0)
761
+ else:
762
+ beta = 0.0
763
+
764
+ kaiser_window = torch.kaiser_window(kernel_size, beta=beta, periodic=False, dtype=torch.float32)
765
+
766
+ # Compute time indices
767
+ if is_even:
768
+ time_indices = torch.arange(-half_size, half_size) + 0.5
769
+ else:
770
+ time_indices = torch.arange(kernel_size) - half_size
771
+
772
+ # Compute sinc filter
773
+ if cutoff == 0:
774
+ return torch.zeros((1, 1, kernel_size), dtype=torch.float32) # Ensures correct shape
775
+
776
+ sinc_filter = torch.sinc(2 * cutoff * time_indices)
777
+ normalized_filter = 2 * cutoff * kaiser_window * sinc_filter
778
+
779
+ # Normalize to ensure sum = 1 (avoid leakage of constant component)
780
+ normalized_filter /= normalized_filter.sum()
781
+
782
+ return normalized_filter.view(1, 1, kernel_size)
783
+
784
+
785
+ class UpSample1d(nn.Module):
786
+ def __init__(self, ratio=2, kernel_size=None):
787
+ super().__init__()
788
+ self.ratio = ratio
789
+ self.kernel_size = int(6 * ratio // 2) * 2 if kernel_size is None else kernel_size
790
+ self.stride = ratio
791
+ self.pad = self.kernel_size // ratio - 1
792
+ self.pad_left = self.pad * self.stride + (self.kernel_size - self.stride) // 2
793
+ self.pad_right = self.pad * self.stride + (self.kernel_size - self.stride + 1) // 2
794
+
795
+ filter = kaiser_sinc_filter1d(cutoff=0.5 / ratio, half_width=0.6 / ratio, kernel_size=self.kernel_size)
796
+ self.register_buffer("filter", filter, persistent=False)
797
+
798
+ def forward(self, hidden_states):
799
+ channels = hidden_states.shape[1]
800
+
801
+ hidden_states = F.pad(hidden_states, (self.pad, self.pad), mode="replicate")
802
+ hidden_states = self.ratio * F.conv_transpose1d(
803
+ hidden_states, self.filter.expand(channels, -1, -1), stride=self.stride, groups=channels
804
+ )
805
+ hidden_states = hidden_states[..., self.pad_left : -self.pad_right]
806
+
807
+ return hidden_states
808
+
809
+
810
+ class DownSample1d(nn.Module):
811
+ def __init__(self, ratio=2, kernel_size=None):
812
+ super().__init__()
813
+ cutoff = 0.5 / ratio
814
+ half_width = 0.6 / ratio
815
+
816
+ if cutoff < 0.0:
817
+ raise ValueError("Minimum cutoff must be larger than zero.")
818
+ if cutoff > 0.5:
819
+ raise ValueError("A cutoff above 0.5 does not make sense.")
820
+
821
+ self.even = kernel_size % 2 == 0
822
+ self.pad_left = kernel_size // 2 - int(self.even)
823
+ self.pad_right = kernel_size // 2
824
+ self.stride = ratio
825
+ filter = kaiser_sinc_filter1d(cutoff, half_width, kernel_size)
826
+ self.register_buffer("filter", filter, persistent=False)
827
+
828
+ def forward(self, hidden_states):
829
+ channels = hidden_states.shape[1]
830
+ hidden_states = F.pad(hidden_states, (self.pad_left, self.pad_right), mode="replicate")
831
+ out = F.conv1d(hidden_states, self.filter.expand(channels, -1, -1), stride=self.stride, groups=channels)
832
+ return out
833
+
834
+
835
+ class TorchActivation1d(nn.Module):
836
+ def __init__(
837
+ self,
838
+ activation,
839
+ up_ratio: int = 2,
840
+ down_ratio: int = 2,
841
+ up_kernel_size: int = 12,
842
+ down_kernel_size: int = 12,
843
+ ):
844
+ super().__init__()
845
+ if not callable(activation):
846
+ raise TypeError("Activation function must be callable")
847
+ self.act = activation
848
+ self.upsample = UpSample1d(up_ratio, up_kernel_size)
849
+ self.downsample = DownSample1d(down_ratio, down_kernel_size)
850
+
851
+ def forward(self, hidden_states):
852
+ hidden_states = self.upsample(hidden_states)
853
+ hidden_states = self.act(hidden_states)
854
+ hidden_states = self.downsample(hidden_states)
855
+
856
+ return hidden_states
857
+
858
+
859
+ class CausalConv1d(nn.Conv1d):
860
+ def __init__(self, *args, **kwargs):
861
+ super().__init__(*args, **kwargs)
862
+ self.causal_padding = self.dilation[0] * (self.kernel_size[0] - 1)
863
+
864
+ def forward(self, x):
865
+ return self._conv_forward(F.pad(x, [self.causal_padding, 0]), self.weight, self.bias)
866
+
867
+
868
+ class AMPBlock(torch.nn.Module):
869
+ def __init__(
870
+ self,
871
+ channels,
872
+ kernel_size=3,
873
+ dilation=(1, 3, 5),
874
+ causal_type='1',
875
+ ):
876
+ super().__init__()
877
+
878
+ self.convs1 = nn.ModuleList(
879
+ [
880
+ CausalConv1d(
881
+ channels,
882
+ channels,
883
+ kernel_size,
884
+ 1,
885
+ dilation=dilation[0],
886
+ ),
887
+ CausalConv1d(
888
+ channels,
889
+ channels,
890
+ kernel_size,
891
+ 1,
892
+ dilation=dilation[1],
893
+ ),
894
+ CausalConv1d(
895
+ channels,
896
+ channels,
897
+ kernel_size,
898
+ 1,
899
+ dilation=dilation[2],
900
+ ),
901
+ ]
902
+ )
903
+
904
+ if causal_type == '1':
905
+ self.convs2 = nn.ModuleList(
906
+ [
907
+ nn.Conv1d(
908
+ channels,
909
+ channels,
910
+ kernel_size,
911
+ 1,
912
+ dilation=1,
913
+ padding=self._get_padding(kernel_size, 1),
914
+ ),
915
+ nn.Conv1d(
916
+ channels,
917
+ channels,
918
+ kernel_size,
919
+ 1,
920
+ dilation=1,
921
+ padding=self._get_padding(kernel_size, 1),
922
+ ),
923
+ nn.Conv1d(
924
+ channels,
925
+ channels,
926
+ kernel_size,
927
+ 1,
928
+ dilation=1,
929
+ padding=self._get_padding(kernel_size, 1),
930
+ ),
931
+ ]
932
+ )
933
+ else:
934
+ self.convs2 = nn.ModuleList(
935
+ [
936
+ CausalConv1d(
937
+ channels,
938
+ channels,
939
+ kernel_size,
940
+ 1,
941
+ dilation=1,
942
+ ),
943
+ CausalConv1d(
944
+ channels,
945
+ channels,
946
+ kernel_size,
947
+ 1,
948
+ dilation=1,
949
+ ),
950
+ CausalConv1d(
951
+ channels,
952
+ channels,
953
+ kernel_size,
954
+ 1,
955
+ dilation=1,
956
+ ),
957
+ ]
958
+ )
959
+
960
+ self.num_layers = len(self.convs1) + len(self.convs2) # total number of conv layers
961
+
962
+ self.activations = nn.ModuleList(
963
+ [TorchActivation1d(activation=SnakeBeta(channels)) for _ in range(self.num_layers)]
964
+ )
965
+
966
+ if causal_type == '2':
967
+ self.pre_conv = nn.Conv1d(
968
+ channels,
969
+ channels,
970
+ kernel_size,
971
+ stride=1,
972
+ padding=self._get_padding(kernel_size, 1),
973
+ )
974
+ self.pre_act = TorchActivation1d(activation=SnakeBeta(channels))
975
+ else:
976
+ self.pre_conv = nn.Identity()
977
+ self.pre_act = nn.Identity()
978
+
979
+ def _get_padding(self, kernel_size, dilation=1):
980
+ return int((kernel_size * dilation - dilation) / 2)
981
+
982
+ def forward(self, x):
983
+ hidden_states = self.pre_conv(x)
984
+ hidden_states = self.pre_act(hidden_states)
985
+ acts1, acts2 = self.activations[::2], self.activations[1::2]
986
+ for conv1, conv2, act1, act2 in zip(self.convs1, self.convs2, acts1, acts2):
987
+ hidden_states = act1(hidden_states)
988
+ hidden_states = conv1(hidden_states)
989
+ hidden_states = act2(hidden_states)
990
+ hidden_states = conv2(hidden_states)
991
+ x = x + hidden_states
992
+ return x
993
+
994
+
995
+ @auto_docstring
996
+ class Qwen3TTSTokenizerV1DecoderBigVGANModel(Qwen3TTSTokenizerV1DecoderPreTrainedModel):
997
+ config: Qwen3TTSTokenizerV1DecoderBigVGANConfig
998
+
999
+ def __init__(self, config: Qwen3TTSTokenizerV1DecoderBigVGANConfig):
1000
+ super().__init__(config)
1001
+ self.num_residual_blocks = len(config.resblock_kernel_sizes)
1002
+ self.num_upsample_layers = len(config.upsample_rates)
1003
+
1004
+ self.conv_pre = nn.Conv1d(config.mel_dim, config.upsample_initial_channel, 5, 1, padding=2)
1005
+
1006
+ # Removing extra ModuleList breaks official state dict
1007
+ ups = [
1008
+ nn.ModuleList(
1009
+ [
1010
+ nn.ConvTranspose1d(
1011
+ config.upsample_initial_channel // (2**layer_idx),
1012
+ config.upsample_initial_channel // (2 ** (layer_idx + 1)),
1013
+ kernel_size,
1014
+ stride,
1015
+ padding=(kernel_size - stride) // 2,
1016
+ )
1017
+ ]
1018
+ )
1019
+ for layer_idx, (stride, kernel_size) in enumerate(zip(config.upsample_rates, config.upsample_kernel_sizes))
1020
+ ]
1021
+ self.ups = nn.ModuleList(ups)
1022
+
1023
+ self.resblocks = nn.ModuleList(
1024
+ [
1025
+ AMPBlock(config.upsample_initial_channel // (2 ** (layer_idx + 1)), kernel_size, dilation, '1' if layer_idx > 1 else '2')
1026
+ for layer_idx in range(self.num_upsample_layers)
1027
+ for kernel_size, dilation in zip(config.resblock_kernel_sizes, config.resblock_dilation_sizes)
1028
+ ]
1029
+ )
1030
+
1031
+ self.activation_post = TorchActivation1d(
1032
+ activation=SnakeBeta(config.upsample_initial_channel // (2**self.num_upsample_layers))
1033
+ )
1034
+ self.conv_post = nn.Conv1d(
1035
+ config.upsample_initial_channel // (2**self.num_upsample_layers), 1, 7, 1, padding=3, bias=False
1036
+ )
1037
+
1038
+ def normalize_spectrogram(self, spectrogram, max_value, min_db):
1039
+ return torch.clamp((2 * max_value) * ((spectrogram - min_db) / (-min_db)) - max_value, -max_value, max_value)
1040
+
1041
+ def amplitude_to_db(self, amplitude, min_db_level):
1042
+ min_level = torch.exp(
1043
+ torch.tensor(min_db_level / 20.0 * np.log(10), device=amplitude.device, dtype=amplitude.dtype)
1044
+ )
1045
+ return 20 * torch.log10(torch.clamp(amplitude, min=min_level))
1046
+
1047
+ def process_mel_spectrogram(self, mel_spectrogram):
1048
+ amplitude_spectrum = torch.exp(mel_spectrogram)
1049
+ decibel_spectrum = self.amplitude_to_db(amplitude_spectrum, -115) - 20
1050
+ return self.normalize_spectrogram(decibel_spectrum, 1, -115)
1051
+
1052
+ def forward(self, mel_spectrogram):
1053
+ processed_spectrogram = self.process_mel_spectrogram(mel_spectrogram)
1054
+ hidden_representation = self.conv_pre(processed_spectrogram)
1055
+
1056
+ for layer_index in range(self.num_upsample_layers):
1057
+ hidden_representation = self.ups[layer_index][0](hidden_representation)
1058
+ residual_output = sum(
1059
+ self.resblocks[layer_index * self.num_residual_blocks + block_index](hidden_representation)
1060
+ for block_index in range(self.num_residual_blocks)
1061
+ )
1062
+ residual_output = residual_output / self.num_residual_blocks
1063
+ hidden_representation = residual_output
1064
+
1065
+ hidden_representation = self.activation_post(hidden_representation)
1066
+ output_waveform = self.conv_post(hidden_representation)
1067
+ return torch.clamp(output_waveform, min=-1.0, max=1.0).squeeze(1)
1068
+
1069
+
1070
+ @auto_docstring
1071
+ class Qwen3TTSTokenizerV1DecoderDiTModel(Qwen3TTSTokenizerV1DecoderPreTrainedModel):
1072
+ config: Qwen3TTSTokenizerV1DecoderDiTConfig
1073
+ _no_split_modules = ["DiTDecoderLayer"]
1074
+
1075
+ def __init__(self, config: Qwen3TTSTokenizerV1DecoderDiTConfig):
1076
+ super().__init__(config)
1077
+ self.mel_dim = config.mel_dim
1078
+ self.repeats = config.repeats
1079
+ self.time_embed = DiTTimestepEmbedding(config.hidden_size)
1080
+
1081
+ self.text_embed = DiTCodecEmbedding(config.num_embeds, config.emb_dim, config.repeats)
1082
+ self.input_embed = DiTInputEmbedding(config)
1083
+
1084
+ self.rotary_embed = Qwen3TTSTokenizerV1DecoderDiTRotaryEmbedding(config.head_dim)
1085
+
1086
+ self.hidden_size = config.hidden_size
1087
+ self.layers = config.num_hidden_layers
1088
+ self.block_size = config.block_size
1089
+ self.num_attention_heads = config.num_attention_heads
1090
+
1091
+ self.transformer_blocks = nn.ModuleList()
1092
+ for i in range(config.num_hidden_layers):
1093
+ self.transformer_blocks.append(
1094
+ DiTDecoderLayer(
1095
+ config,
1096
+ look_ahead_block=1 if i in config.look_ahead_layers else 0,
1097
+ look_backward_block=1 if i in config.look_backward_layers else 0,
1098
+ )
1099
+ )
1100
+
1101
+ self.norm_out = AdaLayerNormZero_Final(config.hidden_size) # final modulation
1102
+ self.proj_out = nn.Linear(config.hidden_size, config.mel_dim)
1103
+
1104
+ def _create_block_diff(self, hidden_states):
1105
+ batch, seq_len = hidden_states.shape[0], hidden_states.shape[1]
1106
+ block_indices = torch.arange(seq_len, device=hidden_states.device) // self.block_size # [seq_length]
1107
+
1108
+ block_i = block_indices.unsqueeze(1) # [seq_length, 1]
1109
+ block_j = block_indices.unsqueeze(0) # [1, seq_length]
1110
+ block_diff = block_j - block_i # (n, n)
1111
+
1112
+ return block_diff.expand(batch, self.num_attention_heads, seq_len, seq_len)
1113
+
1114
+ def forward(
1115
+ self,
1116
+ hidden_states,
1117
+ condition_vector,
1118
+ speaker_embedding,
1119
+ quantized_code,
1120
+ time_step,
1121
+ drop_audio_conditioning=False,
1122
+ drop_code=False,
1123
+ apply_cfg=True,
1124
+ ):
1125
+ batch_size = hidden_states.shape[0] * 2
1126
+ if time_step.ndim == 0:
1127
+ time_step = time_step.repeat(batch_size)
1128
+
1129
+ # Compute embeddings
1130
+ time_embedding = self.time_embed(time_step)
1131
+ text_embedding = self.text_embed(quantized_code, drop_code=False if apply_cfg else drop_code)
1132
+ text_embedding_unconditioned = self.text_embed(quantized_code, drop_code=True) if apply_cfg else None
1133
+
1134
+ hidden_states = self.input_embed(
1135
+ hidden_states,
1136
+ speaker_embedding,
1137
+ condition_vector,
1138
+ text_embedding,
1139
+ drop_audio_cond=drop_audio_conditioning,
1140
+ code_embed_uncond=text_embedding_unconditioned,
1141
+ apply_cfg=apply_cfg,
1142
+ )
1143
+
1144
+ # Compute positional encodings
1145
+ position_embeddings = self.rotary_embed(hidden_states)
1146
+ blockwise_difference = self._create_block_diff(hidden_states)
1147
+
1148
+ # Transformer blocks
1149
+ for transformer_block in self.transformer_blocks:
1150
+ hidden_states = transformer_block(
1151
+ hidden_states,
1152
+ time_embedding,
1153
+ position_embeddings=position_embeddings,
1154
+ block_diff=blockwise_difference,
1155
+ )
1156
+
1157
+ hidden_states = self.norm_out(hidden_states, time_embedding)
1158
+ output = self.proj_out(hidden_states)
1159
+
1160
+ return output
1161
+
1162
+ def optimized_scale(self, positive_flat, negative_flat):
1163
+ # Calculate dot production
1164
+ dot_product = torch.sum(positive_flat * negative_flat, dim=1, keepdim=True)
1165
+ # Squared norm of uncondition
1166
+ squared_norm = torch.sum(negative_flat ** 2, dim=1, keepdim=True) + 1e-8
1167
+ # st_star = v_cond^T * v_uncond / ||v_uncond||^2
1168
+ st_star = dot_product / squared_norm
1169
+ return st_star
1170
+
1171
+ @torch.no_grad()
1172
+ def sample(
1173
+ self,
1174
+ conditioning_vector,
1175
+ reference_mel_spectrogram,
1176
+ quantized_code,
1177
+ num_steps=10,
1178
+ guidance_scale=0.5,
1179
+ sway_coefficient=-1.0,
1180
+ ):
1181
+ noise_initialization = torch.randn([quantized_code.shape[0], 30000, self.mel_dim], dtype=reference_mel_spectrogram.dtype)
1182
+ maximum_duration = quantized_code.shape[1] * self.repeats
1183
+ initial_state = noise_initialization[:, :maximum_duration].to(quantized_code.device)
1184
+ conditioning_vector = conditioning_vector.unsqueeze(1).repeat(1, maximum_duration, 1)
1185
+
1186
+ def ode_function(time_step, hidden_states):
1187
+ if guidance_scale < 1e-5:
1188
+ prediction = self(
1189
+ hidden_states=hidden_states,
1190
+ speaker_embedding=conditioning_vector,
1191
+ condition_vector=reference_mel_spectrogram,
1192
+ quantized_code=quantized_code,
1193
+ time_step=time_step,
1194
+ drop_audio_conditioning=False,
1195
+ drop_code=False,
1196
+ )
1197
+ return prediction
1198
+
1199
+ model_output = self(
1200
+ hidden_states=hidden_states,
1201
+ quantized_code=quantized_code,
1202
+ speaker_embedding=conditioning_vector,
1203
+ condition_vector=reference_mel_spectrogram,
1204
+ time_step=time_step,
1205
+ apply_cfg=True,
1206
+ )
1207
+ guided_prediction, null_prediction = torch.chunk(model_output, 2, dim=0)
1208
+
1209
+ return guided_prediction + (guided_prediction - null_prediction) * guidance_scale
1210
+
1211
+ initial_time = 0
1212
+ time_embedding = torch.linspace(
1213
+ initial_time, 1, num_steps, device=quantized_code.device, dtype=conditioning_vector.dtype
1214
+ )
1215
+
1216
+ if sway_coefficient is not None:
1217
+ time_embedding += sway_coefficient * (torch.cos(torch.pi / 2 * time_embedding) - 1 + time_embedding)
1218
+
1219
+ values = initial_state.clone()
1220
+ for t0, t1 in zip(time_embedding[:-1], time_embedding[1:]):
1221
+ dt = t1 - t0
1222
+ vt = ode_function(t0, values)
1223
+ values = values + vt * dt
1224
+
1225
+ generated_mel_spectrogram = values.permute(0, 2, 1)
1226
+ return generated_mel_spectrogram
1227
+
1228
+
1229
+ @auto_docstring
1230
+ class Qwen3TTSTokenizerV1Decoder(Qwen3TTSTokenizerV1DecoderPreTrainedModel):
1231
+ config: Qwen3TTSTokenizerV1DecoderConfig
1232
+ base_model_prefix = "model"
1233
+ _no_split_modules = ["Qwen3TTSTokenizerV1DecoderDiTModel", "Qwen3TTSTokenizerV1DecoderBigVGANModel"]
1234
+
1235
+ def __init__(self, config: Qwen3TTSTokenizerV1DecoderConfig):
1236
+ super().__init__(config)
1237
+ attn_impl = config._attn_implementation
1238
+ if config._attn_implementation == "flash_attention_2":
1239
+ logger.warning_once(
1240
+ "Qwen3TTSTokenizerV1Decoder must inference with fp32, but flash_attention_2 only supports fp16 and bf16, "
1241
+ "attention implementation of Qwen3TTSTokenizerV1Decoder will fallback to sdpa."
1242
+ )
1243
+ attn_impl = "sdpa"
1244
+ elif config._attn_implementation == "eager":
1245
+ logger.warning_once(
1246
+ "Qwen3TTSTokenizerV1Decoder does not support eager attention implementation, fall back to sdpa"
1247
+ )
1248
+ attn_impl = "sdpa"
1249
+ self.dit = Qwen3TTSTokenizerV1DecoderDiTModel._from_config(
1250
+ config.dit_config, attn_implementation=attn_impl
1251
+ )
1252
+ self.bigvgan = Qwen3TTSTokenizerV1DecoderBigVGANModel._from_config(
1253
+ config.bigvgan_config, attn_implementation=attn_impl
1254
+ )
1255
+
1256
+ def forward(
1257
+ self,
1258
+ code,
1259
+ conditioning,
1260
+ reference_mel,
1261
+ num_steps=10,
1262
+ guidance_scale=0.5,
1263
+ sway_coefficient=-1.0,
1264
+ **kwargs,
1265
+ ):
1266
+ """Generates a waveform from input code and conditioning parameters."""
1267
+
1268
+ mel_spectrogram = self.dit.sample(
1269
+ conditioning,
1270
+ reference_mel,
1271
+ code,
1272
+ num_steps=num_steps,
1273
+ guidance_scale=guidance_scale,
1274
+ sway_coefficient=sway_coefficient,
1275
+ )
1276
+
1277
+ waveform = self.bigvgan(mel_spectrogram)
1278
+
1279
+ return waveform
1280
+
1281
+
1282
+ class Qwen3TTSTokenizerV1Encoder(Qwen3TTSTokenizerV1EncoderPreTrainedModel):
1283
+ config: Qwen3TTSTokenizerV1EncoderConfig
1284
+ def __init__(self, config: Qwen3TTSTokenizerV1EncoderConfig):
1285
+ super().__init__(config)
1286
+
1287
+ self.tokenizer = WhisperEncoderVQ(
1288
+ n_mels=config.n_mels,
1289
+ n_ctx=config.n_ctx,
1290
+ n_state=config.n_state,
1291
+ n_head=config.n_head,
1292
+ n_layer=config.n_layer,
1293
+ n_window=config.n_window,
1294
+ output_dim=config.output_dim,
1295
+ grad_checkpointing=config.grad_checkpointing,
1296
+ enable_mp=config.enable_mp,
1297
+ audio_sequence_parallel=config.audio_sequence_parallel,
1298
+ audio_vq_type=config.audio_vq_type,
1299
+ audio_vq_layers=config.audio_vq_layers,
1300
+ audio_vq_codebook_size=config.audio_vq_codebook_size,
1301
+ audio_vq_codebook_dim=config.audio_vq_codebook_dim,
1302
+ audio_vq_pe=config.audio_vq_pe,
1303
+ audio_vq_ds_rate=config.audio_vq_ds_rate,
1304
+ )
1305
+
1306
+ self.padding = True
1307
+ self.audio_vq_ds_rate = self.tokenizer.audio_vq_ds_rate
1308
+
1309
+ def speech2mel(self, speechs):
1310
+ mels = [
1311
+ get_mel_audio(
1312
+ speech, padding = self.padding, audio_vq_ds_rate = self.audio_vq_ds_rate
1313
+ ).to(speech.dtype).to(self.tokenizer.conv1.weight.device)
1314
+ for speech in speechs
1315
+ ]
1316
+ return mels
1317
+
1318
+ def mel2code(self, mels):
1319
+ audio_mellens = [mel.size(-1) for mel in mels]
1320
+ audio_aftercnnlens = [get_T_after_cnn(T) for T in audio_mellens]
1321
+ audio_seqlens = [T + 2 for T in audio_aftercnnlens]
1322
+
1323
+ with torch.no_grad():
1324
+ _, indices = self.tokenizer(
1325
+ x_list = mels,
1326
+ audio_mellens = audio_mellens,
1327
+ audio_aftercnnlens = audio_aftercnnlens,
1328
+ audio_seqlens = audio_seqlens,
1329
+ return_indices=True,
1330
+ )
1331
+
1332
+ indice_lens = [T // self.tokenizer.audio_vq_ds_rate for T in audio_aftercnnlens]
1333
+ indices = pad_sequence(torch.split(indices, indice_lens), batch_first=True, padding_value=0)
1334
+
1335
+ return indices, indice_lens
1336
+
1337
+ def quantize_speech(self, speechs):
1338
+ mels = self.speech2mel(speechs)
1339
+ indices, indice_lens = self.mel2code(mels)
1340
+ return indices, indice_lens
1341
+
1342
+
1343
+ @auto_docstring
1344
+ class Qwen3TTSTokenizerV1PreTrainedModel(PreTrainedModel):
1345
+ config: Qwen3TTSTokenizerV1Config
1346
+ base_model_prefix = "model"
1347
+ supports_gradient_checkpointing = True
1348
+ _skip_keys_device_placement = "past_key_values"
1349
+ _supports_flash_attn = True
1350
+ _supports_sdpa = True
1351
+ _can_compile_fullgraph = False
1352
+ _supports_attention_backend = True
1353
+
1354
+
1355
+ @auto_docstring(
1356
+ custom_intro="""
1357
+ The Qwen3TTSTokenizerV1 model.
1358
+ """
1359
+ )
1360
+ class Qwen3TTSTokenizerV1Model(Qwen3TTSTokenizerV1PreTrainedModel):
1361
+ def __init__(self, config: Qwen3TTSTokenizerV1Config):
1362
+ super().__init__(config)
1363
+ self.config = config
1364
+
1365
+ self.input_sample_rate = config.input_sample_rate
1366
+ self.output_sample_rate = config.output_sample_rate
1367
+
1368
+ self.decode_upsample_rate = config.decode_upsample_rate
1369
+ self.encode_downsample_rate = config.encode_downsample_rate
1370
+
1371
+ self.encoder = Qwen3TTSTokenizerV1Encoder._from_config(self.config.encoder_config)
1372
+ self.decoder = Qwen3TTSTokenizerV1Decoder._from_config(self.config.decoder_config)
1373
+
1374
+ self.encoder_xvector_extractor = None
1375
+
1376
+ self.post_init()
1377
+
1378
+ def load_encoder_xvector_extractor(self, model_path):
1379
+ self.encoder_xvector_extractor = XVectorExtractor(model_path)
1380
+
1381
+ def get_model_type(self):
1382
+ return self.config.model_type
1383
+
1384
+ def get_input_sample_rate(self):
1385
+ return self.input_sample_rate
1386
+
1387
+ def get_output_sample_rate(self):
1388
+ return self.output_sample_rate
1389
+
1390
+ def get_encode_downsample_rate(self):
1391
+ return self.encode_downsample_rate
1392
+
1393
+ def get_decode_upsample_rate(self):
1394
+ return self.decode_upsample_rate
1395
+
1396
+ @classmethod
1397
+ def from_pretrained(
1398
+ cls,
1399
+ pretrained_model_name_or_path,
1400
+ *model_args,
1401
+ config=None,
1402
+ cache_dir=None,
1403
+ ignore_mismatched_sizes=False,
1404
+ force_download=False,
1405
+ local_files_only=False,
1406
+ token=None,
1407
+ revision="main",
1408
+ use_safetensors=None,
1409
+ weights_only=True,
1410
+ **kwargs,
1411
+ ):
1412
+ model = super().from_pretrained(
1413
+ pretrained_model_name_or_path,
1414
+ *model_args,
1415
+ config=config,
1416
+ cache_dir=cache_dir,
1417
+ ignore_mismatched_sizes=ignore_mismatched_sizes,
1418
+ force_download=force_download,
1419
+ local_files_only=local_files_only,
1420
+ token=token,
1421
+ revision=revision,
1422
+ use_safetensors=use_safetensors,
1423
+ weights_only=weights_only,
1424
+ **kwargs,
1425
+ )
1426
+ encoder_xvector_extractor_path = cached_file(
1427
+ pretrained_model_name_or_path,
1428
+ "campplus.onnx",
1429
+ subfolder=kwargs.pop("subfolder", None),
1430
+ cache_dir=kwargs.pop("cache_dir", None),
1431
+ force_download=kwargs.pop("force_download", False),
1432
+ proxies=kwargs.pop("proxies", None),
1433
+ resume_download=kwargs.pop("resume_download", None),
1434
+ local_files_only=kwargs.pop("local_files_only", False),
1435
+ token=kwargs.pop("use_auth_token", None),
1436
+ revision=kwargs.pop("revision", None),
1437
+ )
1438
+ if encoder_xvector_extractor_path is None:
1439
+ raise ValueError(f"""{pretrained_model_name_or_path}/{encoder_xvector_extractor_path} not exists""")
1440
+ model.load_encoder_xvector_extractor(encoder_xvector_extractor_path)
1441
+
1442
+ return model
1443
+
1444
+ def encode(
1445
+ self,
1446
+ input_values: torch.Tensor,
1447
+ padding_mask: Optional[torch.Tensor] = None,
1448
+ return_dict: Optional[bool] = None,
1449
+ ) -> Union[tuple[torch.Tensor, Optional[torch.Tensor]], Qwen3TTSTokenizerV1EncoderOutput]:
1450
+ """
1451
+ Encodes the input audio waveform into discrete codes.
1452
+
1453
+ Args:
1454
+ input_values (`torch.Tensor` of shape `(batch_size, sequence_length)`):
1455
+ Float values of the input audio waveform.
1456
+ padding_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`):
1457
+ Indicates which inputs are to be ignored due to padding, where elements are either 1 for *not masked* or 0
1458
+ for *masked*.
1459
+ return_dict (`bool`, *optional*):
1460
+ Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
1461
+ """
1462
+ return_dict = return_dict if return_dict is not None else self.config.return_dict
1463
+
1464
+ wavs = [value[:mask.sum()] for value, mask in zip(input_values, padding_mask)]
1465
+
1466
+ codes, codes_lens = self.encoder.quantize_speech(wavs)
1467
+ codes = [c[:l] for c, l in zip(codes, codes_lens)]
1468
+
1469
+ xvectors = []
1470
+ ref_mels = []
1471
+ for wav in wavs:
1472
+ xvector, ref_mel = self.encoder_xvector_extractor.extract_code(wav.cpu().numpy())
1473
+ xvector = torch.tensor(xvector).to(wav.dtype).to(wav.device)
1474
+ ref_mel = torch.tensor(ref_mel).to(wav.dtype).to(wav.device)
1475
+ xvectors.append(xvector)
1476
+ ref_mels.append(ref_mel)
1477
+
1478
+ if not return_dict:
1479
+ return (
1480
+ codes,
1481
+ xvectors,
1482
+ ref_mels
1483
+ )
1484
+
1485
+ return Qwen3TTSTokenizerV1EncoderOutput(codes, xvectors, ref_mels)
1486
+
1487
+ def decode(
1488
+ self,
1489
+ audio_codes: torch.Tensor,
1490
+ xvectors: torch.Tensor,
1491
+ ref_mels: torch.Tensor,
1492
+ return_dict: Optional[bool] = None,
1493
+ ) -> Union[tuple[torch.Tensor, torch.Tensor], Qwen3TTSTokenizerV1DecoderOutput]:
1494
+ """
1495
+ Decodes the given frames into an output audio waveform.
1496
+
1497
+ Note that the output might be a bit bigger than the input. In that case, any extra steps at the end can be
1498
+ trimmed.
1499
+
1500
+ Args:
1501
+ audio_codes (`torch.LongTensor` of shape `(batch_size, codes_length)`, *optional*):
1502
+ Discret code embeddings computed using `model.encode`.
1503
+ xvectors (`torch.FloatTensor` of shape `(batch_size, xvector_dim)`, *optional*):
1504
+ X-vector embeddings computed using `model.encode`.
1505
+ ref_mels (`torch.FloatTensor` of shape `(batch_size, mel_length, mel_dim)`, *optional*):
1506
+ Reference mel spectrogram computed using `model.encode`.
1507
+ return_dict (`bool`, *optional*):
1508
+ Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
1509
+
1510
+ """
1511
+ return_dict = return_dict if return_dict is not None else self.config.return_dict
1512
+ audio_lengths = (audio_codes > -1).sum(1) * self.decode_upsample_rate
1513
+
1514
+ audio_codes = torch.clamp(audio_codes, min=0)
1515
+ audio_values = self.decoder(code=audio_codes,
1516
+ reference_mel=ref_mels,
1517
+ conditioning=xvectors)
1518
+
1519
+ audio_values = [a[:l] for a, l in zip(audio_values, audio_lengths)]
1520
+
1521
+ if not return_dict:
1522
+ return (
1523
+ audio_values,
1524
+ )
1525
+
1526
+ return Qwen3TTSTokenizerV1DecoderOutput(audio_values)
1527
+
1528
+
1529
+ __all__ = ["Qwen3TTSTokenizerV1Model", "Qwen3TTSTokenizerV1PreTrainedModel"]