melottss 0.0.1__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (69) hide show
  1. melo/__init__.py +0 -0
  2. melo/api.py +135 -0
  3. melo/app.py +61 -0
  4. melo/attentions.py +459 -0
  5. melo/commons.py +160 -0
  6. melo/data_utils.py +413 -0
  7. melo/download_utils.py +67 -0
  8. melo/infer.py +25 -0
  9. melo/init_downloads.py +14 -0
  10. melo/losses.py +58 -0
  11. melo/main.py +36 -0
  12. melo/mel_processing.py +174 -0
  13. melo/models.py +1030 -0
  14. melo/modules.py +598 -0
  15. melo/monotonic_align/__init__.py +16 -0
  16. melo/monotonic_align/core.py +46 -0
  17. melo/preprocess_text.py +135 -0
  18. melo/split_utils.py +174 -0
  19. melo/text/__init__.py +35 -0
  20. melo/text/chinese.py +199 -0
  21. melo/text/chinese_bert.py +107 -0
  22. melo/text/chinese_mix.py +253 -0
  23. melo/text/cleaner.py +68 -0
  24. melo/text/cleaner_multiling.py +110 -0
  25. melo/text/cmudict_cache.pickle +0 -0
  26. melo/text/english.py +284 -0
  27. melo/text/english_bert.py +39 -0
  28. melo/text/english_utils/__init__.py +0 -0
  29. melo/text/english_utils/abbreviations.py +35 -0
  30. melo/text/english_utils/number_norm.py +97 -0
  31. melo/text/english_utils/time_norm.py +47 -0
  32. melo/text/es_phonemizer/__init__.py +0 -0
  33. melo/text/es_phonemizer/base.py +140 -0
  34. melo/text/es_phonemizer/cleaner.py +109 -0
  35. melo/text/es_phonemizer/es_symbols.txt +1 -0
  36. melo/text/es_phonemizer/es_to_ipa.py +12 -0
  37. melo/text/es_phonemizer/example_ipa.txt +400 -0
  38. melo/text/es_phonemizer/gruut_wrapper.py +253 -0
  39. melo/text/es_phonemizer/punctuation.py +174 -0
  40. melo/text/es_phonemizer/spanish_symbols.txt +1 -0
  41. melo/text/fr_phonemizer/__init__.py +0 -0
  42. melo/text/fr_phonemizer/base.py +140 -0
  43. melo/text/fr_phonemizer/cleaner.py +122 -0
  44. melo/text/fr_phonemizer/example_ipa.txt +1 -0
  45. melo/text/fr_phonemizer/fr_to_ipa.py +30 -0
  46. melo/text/fr_phonemizer/french_abbreviations.py +48 -0
  47. melo/text/fr_phonemizer/french_symbols.txt +1 -0
  48. melo/text/fr_phonemizer/gruut_wrapper.py +258 -0
  49. melo/text/fr_phonemizer/punctuation.py +172 -0
  50. melo/text/french.py +94 -0
  51. melo/text/french_bert.py +39 -0
  52. melo/text/japanese.py +662 -0
  53. melo/text/japanese_bert.py +49 -0
  54. melo/text/ko_dictionary.py +44 -0
  55. melo/text/korean.py +192 -0
  56. melo/text/opencpop-strict.txt +429 -0
  57. melo/text/spanish.py +122 -0
  58. melo/text/spanish_bert.py +39 -0
  59. melo/text/symbols.py +290 -0
  60. melo/text/tone_sandhi.py +769 -0
  61. melo/train.py +635 -0
  62. melo/transforms.py +209 -0
  63. melo/utils.py +424 -0
  64. melottss-0.0.1.dist-info/METADATA +5 -0
  65. melottss-0.0.1.dist-info/RECORD +69 -0
  66. melottss-0.0.1.dist-info/WHEEL +5 -0
  67. melottss-0.0.1.dist-info/entry_points.txt +4 -0
  68. melottss-0.0.1.dist-info/licenses/LICENSE +19 -0
  69. melottss-0.0.1.dist-info/top_level.txt +1 -0
melo/__init__.py ADDED
File without changes
melo/api.py ADDED
@@ -0,0 +1,135 @@
1
+ import os
2
+ import re
3
+ import json
4
+ import torch
5
+ import librosa
6
+ import soundfile
7
+ import torchaudio
8
+ import numpy as np
9
+ import torch.nn as nn
10
+ from tqdm import tqdm
11
+ import torch
12
+
13
+ from . import utils
14
+ from . import commons
15
+ from .models import SynthesizerTrn
16
+ from .split_utils import split_sentence
17
+ from .mel_processing import spectrogram_torch, spectrogram_torch_conv
18
+ from .download_utils import load_or_download_config, load_or_download_model
19
+
20
+ class TTS(nn.Module):
21
+ def __init__(self,
22
+ language,
23
+ device='auto',
24
+ use_hf=True,
25
+ config_path=None,
26
+ ckpt_path=None):
27
+ super().__init__()
28
+ if device == 'auto':
29
+ device = 'cpu'
30
+ if torch.cuda.is_available(): device = 'cuda'
31
+ if torch.backends.mps.is_available(): device = 'mps'
32
+ if 'cuda' in device:
33
+ assert torch.cuda.is_available()
34
+
35
+ # config_path =
36
+ hps = load_or_download_config(language, use_hf=use_hf, config_path=config_path)
37
+
38
+ num_languages = hps.num_languages
39
+ num_tones = hps.num_tones
40
+ symbols = hps.symbols
41
+
42
+ model = SynthesizerTrn(
43
+ len(symbols),
44
+ hps.data.filter_length // 2 + 1,
45
+ hps.train.segment_size // hps.data.hop_length,
46
+ n_speakers=hps.data.n_speakers,
47
+ num_tones=num_tones,
48
+ num_languages=num_languages,
49
+ **hps.model,
50
+ ).to(device)
51
+
52
+ model.eval()
53
+ self.model = model
54
+ self.symbol_to_id = {s: i for i, s in enumerate(symbols)}
55
+ self.hps = hps
56
+ self.device = device
57
+
58
+ # load state_dict
59
+ checkpoint_dict = load_or_download_model(language, device, use_hf=use_hf, ckpt_path=ckpt_path)
60
+ self.model.load_state_dict(checkpoint_dict['model'], strict=True)
61
+
62
+ language = language.split('_')[0]
63
+ self.language = 'ZH_MIX_EN' if language == 'ZH' else language # we support a ZH_MIX_EN model
64
+
65
+ @staticmethod
66
+ def audio_numpy_concat(segment_data_list, sr, speed=1.):
67
+ audio_segments = []
68
+ for segment_data in segment_data_list:
69
+ audio_segments += segment_data.reshape(-1).tolist()
70
+ audio_segments += [0] * int((sr * 0.05) / speed)
71
+ audio_segments = np.array(audio_segments).astype(np.float32)
72
+ return audio_segments
73
+
74
+ @staticmethod
75
+ def split_sentences_into_pieces(text, language, quiet=False):
76
+ texts = split_sentence(text, language_str=language)
77
+ if not quiet:
78
+ print(" > Text split to sentences.")
79
+ print('\n'.join(texts))
80
+ print(" > ===========================")
81
+ return texts
82
+
83
+ def tts_to_file(self, text, speaker_id, output_path=None, sdp_ratio=0.2, noise_scale=0.6, noise_scale_w=0.8, speed=1.0, pbar=None, format=None, position=None, quiet=False,):
84
+ language = self.language
85
+ texts = self.split_sentences_into_pieces(text, language, quiet)
86
+ audio_list = []
87
+ if pbar:
88
+ tx = pbar(texts)
89
+ else:
90
+ if position:
91
+ tx = tqdm(texts, position=position)
92
+ elif quiet:
93
+ tx = texts
94
+ else:
95
+ tx = tqdm(texts)
96
+ for t in tx:
97
+ if language in ['EN', 'ZH_MIX_EN']:
98
+ t = re.sub(r'([a-z])([A-Z])', r'\1 \2', t)
99
+ device = self.device
100
+ bert, ja_bert, phones, tones, lang_ids = utils.get_text_for_tts_infer(t, language, self.hps, device, self.symbol_to_id)
101
+ with torch.no_grad():
102
+ x_tst = phones.to(device).unsqueeze(0)
103
+ tones = tones.to(device).unsqueeze(0)
104
+ lang_ids = lang_ids.to(device).unsqueeze(0)
105
+ bert = bert.to(device).unsqueeze(0)
106
+ ja_bert = ja_bert.to(device).unsqueeze(0)
107
+ x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device)
108
+ del phones
109
+ speakers = torch.LongTensor([speaker_id]).to(device)
110
+ audio = self.model.infer(
111
+ x_tst,
112
+ x_tst_lengths,
113
+ speakers,
114
+ tones,
115
+ lang_ids,
116
+ bert,
117
+ ja_bert,
118
+ sdp_ratio=sdp_ratio,
119
+ noise_scale=noise_scale,
120
+ noise_scale_w=noise_scale_w,
121
+ length_scale=1. / speed,
122
+ )[0][0, 0].data.cpu().float().numpy()
123
+ del x_tst, tones, lang_ids, bert, ja_bert, x_tst_lengths, speakers
124
+ #
125
+ audio_list.append(audio)
126
+ torch.cuda.empty_cache()
127
+ audio = self.audio_numpy_concat(audio_list, sr=self.hps.data.sampling_rate, speed=speed)
128
+
129
+ if output_path is None:
130
+ return audio
131
+ else:
132
+ if format:
133
+ soundfile.write(output_path, audio, self.hps.data.sampling_rate, format=format)
134
+ else:
135
+ soundfile.write(output_path, audio, self.hps.data.sampling_rate)
melo/app.py ADDED
@@ -0,0 +1,61 @@
1
+ # WebUI by mrfakename <X @realmrfakename / HF @mrfakename>
2
+ # Demo also available on HF Spaces: https://huggingface.co/spaces/mrfakename/MeloTTS
3
+ import gradio as gr
4
+ import os, torch, io
5
+ # os.system('python -m unidic download')
6
+ print("Make sure you've downloaded unidic (python -m unidic download) for this WebUI to work.")
7
+ from melo.api import TTS
8
+ speed = 1.0
9
+ import tempfile
10
+ import click
11
+ device = 'auto'
12
+ models = {
13
+ 'EN': TTS(language='EN', device=device),
14
+ 'ES': TTS(language='ES', device=device),
15
+ 'FR': TTS(language='FR', device=device),
16
+ 'ZH': TTS(language='ZH', device=device),
17
+ 'JP': TTS(language='JP', device=device),
18
+ 'KR': TTS(language='KR', device=device),
19
+ }
20
+ speaker_ids = models['EN'].hps.data.spk2id
21
+
22
+ default_text_dict = {
23
+ 'EN': 'The field of text-to-speech has seen rapid development recently.',
24
+ 'ES': 'El campo de la conversión de texto a voz ha experimentado un rápido desarrollo recientemente.',
25
+ 'FR': 'Le domaine de la synthèse vocale a connu un développement rapide récemment',
26
+ 'ZH': 'text-to-speech 领域近年来发展迅速',
27
+ 'JP': 'テキスト読み上げの分野は最近急速な発展を遂げています',
28
+ 'KR': '최근 텍스트 음성 변환 분야가 급속도로 발전하고 있습니다.',
29
+ }
30
+
31
+ def synthesize(speaker, text, speed, language, progress=gr.Progress()):
32
+ bio = io.BytesIO()
33
+ models[language].tts_to_file(text, models[language].hps.data.spk2id[speaker], bio, speed=speed, pbar=progress.tqdm, format='wav')
34
+ return bio.getvalue()
35
+ def load_speakers(language, text):
36
+ if text in list(default_text_dict.values()):
37
+ newtext = default_text_dict[language]
38
+ else:
39
+ newtext = text
40
+ return gr.update(value=list(models[language].hps.data.spk2id.keys())[0], choices=list(models[language].hps.data.spk2id.keys())), newtext
41
+ with gr.Blocks() as demo:
42
+ gr.Markdown('# MeloTTS WebUI\n\nA WebUI for MeloTTS.')
43
+ with gr.Group():
44
+ speaker = gr.Dropdown(speaker_ids.keys(), interactive=True, value='EN-US', label='Speaker')
45
+ language = gr.Radio(['EN', 'ES', 'FR', 'ZH', 'JP', 'KR'], label='Language', value='EN')
46
+ speed = gr.Slider(label='Speed', minimum=0.1, maximum=10.0, value=1.0, interactive=True, step=0.1)
47
+ text = gr.Textbox(label="Text to speak", value=default_text_dict['EN'])
48
+ language.input(load_speakers, inputs=[language, text], outputs=[speaker, text])
49
+ btn = gr.Button('Synthesize', variant='primary')
50
+ aud = gr.Audio(interactive=False)
51
+ btn.click(synthesize, inputs=[speaker, text, speed, language], outputs=[aud])
52
+ gr.Markdown('WebUI by [mrfakename](https://twitter.com/realmrfakename).')
53
+ @click.command()
54
+ @click.option('--share', '-s', is_flag=True, show_default=True, default=False, help="Expose a publicly-accessible shared Gradio link usable by anyone with the link. Only share the link with people you trust.")
55
+ @click.option('--host', '-h', default=None)
56
+ @click.option('--port', '-p', type=int, default=None)
57
+ def main(share, host, port):
58
+ demo.queue(api_open=False).launch(show_api=False, share=share, server_name=host, server_port=port)
59
+
60
+ if __name__ == "__main__":
61
+ main()
melo/attentions.py ADDED
@@ -0,0 +1,459 @@
1
+ import math
2
+ import torch
3
+ from torch import nn
4
+ from torch.nn import functional as F
5
+
6
+ from . import commons
7
+ import logging
8
+
9
+ logger = logging.getLogger(__name__)
10
+
11
+
12
+ class LayerNorm(nn.Module):
13
+ def __init__(self, channels, eps=1e-5):
14
+ super().__init__()
15
+ self.channels = channels
16
+ self.eps = eps
17
+
18
+ self.gamma = nn.Parameter(torch.ones(channels))
19
+ self.beta = nn.Parameter(torch.zeros(channels))
20
+
21
+ def forward(self, x):
22
+ x = x.transpose(1, -1)
23
+ x = F.layer_norm(x, (self.channels,), self.gamma, self.beta, self.eps)
24
+ return x.transpose(1, -1)
25
+
26
+
27
+ @torch.jit.script
28
+ def fused_add_tanh_sigmoid_multiply(input_a, input_b, n_channels):
29
+ n_channels_int = n_channels[0]
30
+ in_act = input_a + input_b
31
+ t_act = torch.tanh(in_act[:, :n_channels_int, :])
32
+ s_act = torch.sigmoid(in_act[:, n_channels_int:, :])
33
+ acts = t_act * s_act
34
+ return acts
35
+
36
+
37
+ class Encoder(nn.Module):
38
+ def __init__(
39
+ self,
40
+ hidden_channels,
41
+ filter_channels,
42
+ n_heads,
43
+ n_layers,
44
+ kernel_size=1,
45
+ p_dropout=0.0,
46
+ window_size=4,
47
+ isflow=True,
48
+ **kwargs
49
+ ):
50
+ super().__init__()
51
+ self.hidden_channels = hidden_channels
52
+ self.filter_channels = filter_channels
53
+ self.n_heads = n_heads
54
+ self.n_layers = n_layers
55
+ self.kernel_size = kernel_size
56
+ self.p_dropout = p_dropout
57
+ self.window_size = window_size
58
+
59
+ self.cond_layer_idx = self.n_layers
60
+ if "gin_channels" in kwargs:
61
+ self.gin_channels = kwargs["gin_channels"]
62
+ if self.gin_channels != 0:
63
+ self.spk_emb_linear = nn.Linear(self.gin_channels, self.hidden_channels)
64
+ self.cond_layer_idx = (
65
+ kwargs["cond_layer_idx"] if "cond_layer_idx" in kwargs else 2
66
+ )
67
+ assert (
68
+ self.cond_layer_idx < self.n_layers
69
+ ), "cond_layer_idx should be less than n_layers"
70
+ self.drop = nn.Dropout(p_dropout)
71
+ self.attn_layers = nn.ModuleList()
72
+ self.norm_layers_1 = nn.ModuleList()
73
+ self.ffn_layers = nn.ModuleList()
74
+ self.norm_layers_2 = nn.ModuleList()
75
+
76
+ for i in range(self.n_layers):
77
+ self.attn_layers.append(
78
+ MultiHeadAttention(
79
+ hidden_channels,
80
+ hidden_channels,
81
+ n_heads,
82
+ p_dropout=p_dropout,
83
+ window_size=window_size,
84
+ )
85
+ )
86
+ self.norm_layers_1.append(LayerNorm(hidden_channels))
87
+ self.ffn_layers.append(
88
+ FFN(
89
+ hidden_channels,
90
+ hidden_channels,
91
+ filter_channels,
92
+ kernel_size,
93
+ p_dropout=p_dropout,
94
+ )
95
+ )
96
+ self.norm_layers_2.append(LayerNorm(hidden_channels))
97
+
98
+ def forward(self, x, x_mask, g=None):
99
+ attn_mask = x_mask.unsqueeze(2) * x_mask.unsqueeze(-1)
100
+ x = x * x_mask
101
+ for i in range(self.n_layers):
102
+ if i == self.cond_layer_idx and g is not None:
103
+ g = self.spk_emb_linear(g.transpose(1, 2))
104
+ g = g.transpose(1, 2)
105
+ x = x + g
106
+ x = x * x_mask
107
+ y = self.attn_layers[i](x, x, attn_mask)
108
+ y = self.drop(y)
109
+ x = self.norm_layers_1[i](x + y)
110
+
111
+ y = self.ffn_layers[i](x, x_mask)
112
+ y = self.drop(y)
113
+ x = self.norm_layers_2[i](x + y)
114
+ x = x * x_mask
115
+ return x
116
+
117
+
118
+ class Decoder(nn.Module):
119
+ def __init__(
120
+ self,
121
+ hidden_channels,
122
+ filter_channels,
123
+ n_heads,
124
+ n_layers,
125
+ kernel_size=1,
126
+ p_dropout=0.0,
127
+ proximal_bias=False,
128
+ proximal_init=True,
129
+ **kwargs
130
+ ):
131
+ super().__init__()
132
+ self.hidden_channels = hidden_channels
133
+ self.filter_channels = filter_channels
134
+ self.n_heads = n_heads
135
+ self.n_layers = n_layers
136
+ self.kernel_size = kernel_size
137
+ self.p_dropout = p_dropout
138
+ self.proximal_bias = proximal_bias
139
+ self.proximal_init = proximal_init
140
+
141
+ self.drop = nn.Dropout(p_dropout)
142
+ self.self_attn_layers = nn.ModuleList()
143
+ self.norm_layers_0 = nn.ModuleList()
144
+ self.encdec_attn_layers = nn.ModuleList()
145
+ self.norm_layers_1 = nn.ModuleList()
146
+ self.ffn_layers = nn.ModuleList()
147
+ self.norm_layers_2 = nn.ModuleList()
148
+ for i in range(self.n_layers):
149
+ self.self_attn_layers.append(
150
+ MultiHeadAttention(
151
+ hidden_channels,
152
+ hidden_channels,
153
+ n_heads,
154
+ p_dropout=p_dropout,
155
+ proximal_bias=proximal_bias,
156
+ proximal_init=proximal_init,
157
+ )
158
+ )
159
+ self.norm_layers_0.append(LayerNorm(hidden_channels))
160
+ self.encdec_attn_layers.append(
161
+ MultiHeadAttention(
162
+ hidden_channels, hidden_channels, n_heads, p_dropout=p_dropout
163
+ )
164
+ )
165
+ self.norm_layers_1.append(LayerNorm(hidden_channels))
166
+ self.ffn_layers.append(
167
+ FFN(
168
+ hidden_channels,
169
+ hidden_channels,
170
+ filter_channels,
171
+ kernel_size,
172
+ p_dropout=p_dropout,
173
+ causal=True,
174
+ )
175
+ )
176
+ self.norm_layers_2.append(LayerNorm(hidden_channels))
177
+
178
+ def forward(self, x, x_mask, h, h_mask):
179
+ """
180
+ x: decoder input
181
+ h: encoder output
182
+ """
183
+ self_attn_mask = commons.subsequent_mask(x_mask.size(2)).to(
184
+ device=x.device, dtype=x.dtype
185
+ )
186
+ encdec_attn_mask = h_mask.unsqueeze(2) * x_mask.unsqueeze(-1)
187
+ x = x * x_mask
188
+ for i in range(self.n_layers):
189
+ y = self.self_attn_layers[i](x, x, self_attn_mask)
190
+ y = self.drop(y)
191
+ x = self.norm_layers_0[i](x + y)
192
+
193
+ y = self.encdec_attn_layers[i](x, h, encdec_attn_mask)
194
+ y = self.drop(y)
195
+ x = self.norm_layers_1[i](x + y)
196
+
197
+ y = self.ffn_layers[i](x, x_mask)
198
+ y = self.drop(y)
199
+ x = self.norm_layers_2[i](x + y)
200
+ x = x * x_mask
201
+ return x
202
+
203
+
204
+ class MultiHeadAttention(nn.Module):
205
+ def __init__(
206
+ self,
207
+ channels,
208
+ out_channels,
209
+ n_heads,
210
+ p_dropout=0.0,
211
+ window_size=None,
212
+ heads_share=True,
213
+ block_length=None,
214
+ proximal_bias=False,
215
+ proximal_init=False,
216
+ ):
217
+ super().__init__()
218
+ assert channels % n_heads == 0
219
+
220
+ self.channels = channels
221
+ self.out_channels = out_channels
222
+ self.n_heads = n_heads
223
+ self.p_dropout = p_dropout
224
+ self.window_size = window_size
225
+ self.heads_share = heads_share
226
+ self.block_length = block_length
227
+ self.proximal_bias = proximal_bias
228
+ self.proximal_init = proximal_init
229
+ self.attn = None
230
+
231
+ self.k_channels = channels // n_heads
232
+ self.conv_q = nn.Conv1d(channels, channels, 1)
233
+ self.conv_k = nn.Conv1d(channels, channels, 1)
234
+ self.conv_v = nn.Conv1d(channels, channels, 1)
235
+ self.conv_o = nn.Conv1d(channels, out_channels, 1)
236
+ self.drop = nn.Dropout(p_dropout)
237
+
238
+ if window_size is not None:
239
+ n_heads_rel = 1 if heads_share else n_heads
240
+ rel_stddev = self.k_channels**-0.5
241
+ self.emb_rel_k = nn.Parameter(
242
+ torch.randn(n_heads_rel, window_size * 2 + 1, self.k_channels)
243
+ * rel_stddev
244
+ )
245
+ self.emb_rel_v = nn.Parameter(
246
+ torch.randn(n_heads_rel, window_size * 2 + 1, self.k_channels)
247
+ * rel_stddev
248
+ )
249
+
250
+ nn.init.xavier_uniform_(self.conv_q.weight)
251
+ nn.init.xavier_uniform_(self.conv_k.weight)
252
+ nn.init.xavier_uniform_(self.conv_v.weight)
253
+ if proximal_init:
254
+ with torch.no_grad():
255
+ self.conv_k.weight.copy_(self.conv_q.weight)
256
+ self.conv_k.bias.copy_(self.conv_q.bias)
257
+
258
+ def forward(self, x, c, attn_mask=None):
259
+ q = self.conv_q(x)
260
+ k = self.conv_k(c)
261
+ v = self.conv_v(c)
262
+
263
+ x, self.attn = self.attention(q, k, v, mask=attn_mask)
264
+
265
+ x = self.conv_o(x)
266
+ return x
267
+
268
+ def attention(self, query, key, value, mask=None):
269
+ # reshape [b, d, t] -> [b, n_h, t, d_k]
270
+ b, d, t_s, t_t = (*key.size(), query.size(2))
271
+ query = query.view(b, self.n_heads, self.k_channels, t_t).transpose(2, 3)
272
+ key = key.view(b, self.n_heads, self.k_channels, t_s).transpose(2, 3)
273
+ value = value.view(b, self.n_heads, self.k_channels, t_s).transpose(2, 3)
274
+
275
+ scores = torch.matmul(query / math.sqrt(self.k_channels), key.transpose(-2, -1))
276
+ if self.window_size is not None:
277
+ assert (
278
+ t_s == t_t
279
+ ), "Relative attention is only available for self-attention."
280
+ key_relative_embeddings = self._get_relative_embeddings(self.emb_rel_k, t_s)
281
+ rel_logits = self._matmul_with_relative_keys(
282
+ query / math.sqrt(self.k_channels), key_relative_embeddings
283
+ )
284
+ scores_local = self._relative_position_to_absolute_position(rel_logits)
285
+ scores = scores + scores_local
286
+ if self.proximal_bias:
287
+ assert t_s == t_t, "Proximal bias is only available for self-attention."
288
+ scores = scores + self._attention_bias_proximal(t_s).to(
289
+ device=scores.device, dtype=scores.dtype
290
+ )
291
+ if mask is not None:
292
+ scores = scores.masked_fill(mask == 0, -1e4)
293
+ if self.block_length is not None:
294
+ assert (
295
+ t_s == t_t
296
+ ), "Local attention is only available for self-attention."
297
+ block_mask = (
298
+ torch.ones_like(scores)
299
+ .triu(-self.block_length)
300
+ .tril(self.block_length)
301
+ )
302
+ scores = scores.masked_fill(block_mask == 0, -1e4)
303
+ p_attn = F.softmax(scores, dim=-1) # [b, n_h, t_t, t_s]
304
+ p_attn = self.drop(p_attn)
305
+ output = torch.matmul(p_attn, value)
306
+ if self.window_size is not None:
307
+ relative_weights = self._absolute_position_to_relative_position(p_attn)
308
+ value_relative_embeddings = self._get_relative_embeddings(
309
+ self.emb_rel_v, t_s
310
+ )
311
+ output = output + self._matmul_with_relative_values(
312
+ relative_weights, value_relative_embeddings
313
+ )
314
+ output = (
315
+ output.transpose(2, 3).contiguous().view(b, d, t_t)
316
+ ) # [b, n_h, t_t, d_k] -> [b, d, t_t]
317
+ return output, p_attn
318
+
319
+ def _matmul_with_relative_values(self, x, y):
320
+ """
321
+ x: [b, h, l, m]
322
+ y: [h or 1, m, d]
323
+ ret: [b, h, l, d]
324
+ """
325
+ ret = torch.matmul(x, y.unsqueeze(0))
326
+ return ret
327
+
328
+ def _matmul_with_relative_keys(self, x, y):
329
+ """
330
+ x: [b, h, l, d]
331
+ y: [h or 1, m, d]
332
+ ret: [b, h, l, m]
333
+ """
334
+ ret = torch.matmul(x, y.unsqueeze(0).transpose(-2, -1))
335
+ return ret
336
+
337
+ def _get_relative_embeddings(self, relative_embeddings, length):
338
+ 2 * self.window_size + 1
339
+ # Pad first before slice to avoid using cond ops.
340
+ pad_length = max(length - (self.window_size + 1), 0)
341
+ slice_start_position = max((self.window_size + 1) - length, 0)
342
+ slice_end_position = slice_start_position + 2 * length - 1
343
+ if pad_length > 0:
344
+ padded_relative_embeddings = F.pad(
345
+ relative_embeddings,
346
+ commons.convert_pad_shape([[0, 0], [pad_length, pad_length], [0, 0]]),
347
+ )
348
+ else:
349
+ padded_relative_embeddings = relative_embeddings
350
+ used_relative_embeddings = padded_relative_embeddings[
351
+ :, slice_start_position:slice_end_position
352
+ ]
353
+ return used_relative_embeddings
354
+
355
+ def _relative_position_to_absolute_position(self, x):
356
+ """
357
+ x: [b, h, l, 2*l-1]
358
+ ret: [b, h, l, l]
359
+ """
360
+ batch, heads, length, _ = x.size()
361
+ # Concat columns of pad to shift from relative to absolute indexing.
362
+ x = F.pad(x, commons.convert_pad_shape([[0, 0], [0, 0], [0, 0], [0, 1]]))
363
+
364
+ # Concat extra elements so to add up to shape (len+1, 2*len-1).
365
+ x_flat = x.view([batch, heads, length * 2 * length])
366
+ x_flat = F.pad(
367
+ x_flat, commons.convert_pad_shape([[0, 0], [0, 0], [0, length - 1]])
368
+ )
369
+
370
+ # Reshape and slice out the padded elements.
371
+ x_final = x_flat.view([batch, heads, length + 1, 2 * length - 1])[
372
+ :, :, :length, length - 1 :
373
+ ]
374
+ return x_final
375
+
376
+ def _absolute_position_to_relative_position(self, x):
377
+ """
378
+ x: [b, h, l, l]
379
+ ret: [b, h, l, 2*l-1]
380
+ """
381
+ batch, heads, length, _ = x.size()
382
+ # pad along column
383
+ x = F.pad(
384
+ x, commons.convert_pad_shape([[0, 0], [0, 0], [0, 0], [0, length - 1]])
385
+ )
386
+ x_flat = x.view([batch, heads, length**2 + length * (length - 1)])
387
+ # add 0's in the beginning that will skew the elements after reshape
388
+ x_flat = F.pad(x_flat, commons.convert_pad_shape([[0, 0], [0, 0], [length, 0]]))
389
+ x_final = x_flat.view([batch, heads, length, 2 * length])[:, :, :, 1:]
390
+ return x_final
391
+
392
+ def _attention_bias_proximal(self, length):
393
+ """Bias for self-attention to encourage attention to close positions.
394
+ Args:
395
+ length: an integer scalar.
396
+ Returns:
397
+ a Tensor with shape [1, 1, length, length]
398
+ """
399
+ r = torch.arange(length, dtype=torch.float32)
400
+ diff = torch.unsqueeze(r, 0) - torch.unsqueeze(r, 1)
401
+ return torch.unsqueeze(torch.unsqueeze(-torch.log1p(torch.abs(diff)), 0), 0)
402
+
403
+
404
+ class FFN(nn.Module):
405
+ def __init__(
406
+ self,
407
+ in_channels,
408
+ out_channels,
409
+ filter_channels,
410
+ kernel_size,
411
+ p_dropout=0.0,
412
+ activation=None,
413
+ causal=False,
414
+ ):
415
+ super().__init__()
416
+ self.in_channels = in_channels
417
+ self.out_channels = out_channels
418
+ self.filter_channels = filter_channels
419
+ self.kernel_size = kernel_size
420
+ self.p_dropout = p_dropout
421
+ self.activation = activation
422
+ self.causal = causal
423
+
424
+ if causal:
425
+ self.padding = self._causal_padding
426
+ else:
427
+ self.padding = self._same_padding
428
+
429
+ self.conv_1 = nn.Conv1d(in_channels, filter_channels, kernel_size)
430
+ self.conv_2 = nn.Conv1d(filter_channels, out_channels, kernel_size)
431
+ self.drop = nn.Dropout(p_dropout)
432
+
433
+ def forward(self, x, x_mask):
434
+ x = self.conv_1(self.padding(x * x_mask))
435
+ if self.activation == "gelu":
436
+ x = x * torch.sigmoid(1.702 * x)
437
+ else:
438
+ x = torch.relu(x)
439
+ x = self.drop(x)
440
+ x = self.conv_2(self.padding(x * x_mask))
441
+ return x * x_mask
442
+
443
+ def _causal_padding(self, x):
444
+ if self.kernel_size == 1:
445
+ return x
446
+ pad_l = self.kernel_size - 1
447
+ pad_r = 0
448
+ padding = [[0, 0], [0, 0], [pad_l, pad_r]]
449
+ x = F.pad(x, commons.convert_pad_shape(padding))
450
+ return x
451
+
452
+ def _same_padding(self, x):
453
+ if self.kernel_size == 1:
454
+ return x
455
+ pad_l = (self.kernel_size - 1) // 2
456
+ pad_r = self.kernel_size // 2
457
+ padding = [[0, 0], [0, 0], [pad_l, pad_r]]
458
+ x = F.pad(x, commons.convert_pad_shape(padding))
459
+ return x