diff --git a/paddlespeech/__init__.py b/paddlespeech/__init__.py index 6c7e75c1f..fcf14c111 100644 --- a/paddlespeech/__init__.py +++ b/paddlespeech/__init__.py @@ -13,3 +13,7 @@ # limitations under the License. import _locale _locale._getdefaultlocale = (lambda *args: ['en_US', 'utf8']) + + + + diff --git a/paddlespeech/audio/functional/mel_extract.py b/paddlespeech/audio/functional/mel_extract.py new file mode 100644 index 000000000..11d881e05 --- /dev/null +++ b/paddlespeech/audio/functional/mel_extract.py @@ -0,0 +1,84 @@ +import numpy as np +import paddle +from librosa.filters import mel as librosa_mel_fn +from scipy.io.wavfile import read + +MAX_WAV_VALUE = 32768.0 + + +def load_wav(full_path): + sampling_rate, data = read(full_path) + return data, sampling_rate + + +def dynamic_range_compression(x, C=1, clip_val=1e-05): + return np.log(np.clip(x, a_min=clip_val, a_max=None) * C) + + +def dynamic_range_decompression(x, C=1): + return np.exp(x) / C + + +def dynamic_range_compression_torch(x, C=1, clip_val=1e-05): + return paddle.log(paddle.clamp(x, min=clip_val) * C) + + +def dynamic_range_decompression_torch(x, C=1): + return paddle.exp(x=x) / C + + +def spectral_normalize_torch(magnitudes): + output = dynamic_range_compression_torch(magnitudes) + return output + + +def spectral_de_normalize_torch(magnitudes): + output = dynamic_range_decompression_torch(magnitudes) + return output + + +mel_basis = {} +hann_window = {} + + +def mel_spectrogram( + y, n_fft, num_mels, sampling_rate, hop_size, win_size, fmin, fmax, center=False +): + if paddle.compat.min(y) < -1.0: + print("min value is ", paddle.compat.min(y)) + if paddle.compat.max(y) > 1.0: + print("max value is ", paddle.compat.max(y)) + global mel_basis, hann_window + if f"{str(fmax)}_{str(y.place)}" not in mel_basis: + mel = librosa_mel_fn( + sr=sampling_rate, n_fft=n_fft, n_mels=num_mels, fmin=fmin, fmax=fmax + ) + mel_basis[str(fmax) + "_" + str(y.place)] = ( + paddle.from_numpy(mel).float().to(y.place) + ) + hann_window[str(y.place)] = paddle.audio.functional.get_window( + win_length=win_size, dtype="float32", window="hann" + ).to(y.place) + y = paddle.compat.pad( + y.unsqueeze(1), + (int((n_fft - hop_size) / 2), int((n_fft - hop_size) / 2)), + mode="reflect", + ) + y = y.squeeze(1) + spec = paddle.view_as_real( + paddle.signal.stft( + x=y, + n_fft=n_fft, + hop_length=hop_size, + win_length=win_size, + window=hann_window[str(y.place)], + center=center, + pad_mode="reflect", + normalized=False, + onesided=True, + ) + ) + spec = paddle.sqrt(spec.pow(2).sum(-1) + 1e-09) + spec = paddle.matmul(mel_basis[str(fmax) + "_" + str(y.place)], spec) + spec = spectral_normalize_torch(spec) + return spec diff --git a/paddlespeech/t2s/models/CosyVoice/class_utils.py b/paddlespeech/t2s/models/CosyVoice/class_utils.py new file mode 100644 index 000000000..61d503c10 --- /dev/null +++ b/paddlespeech/t2s/models/CosyVoice/class_utils.py @@ -0,0 +1,40 @@ +import paddle + +# from cosyvoice.cli.model import CosyVoice2Model, CosyVoiceModel +# from cosyvoice.flow.flow import CausalMaskedDiffWithXvec, MaskedDiffWithXvec +# from cosyvoice.hifigan.generator import HiFTGenerator +# from cosyvoice.llm.llm import Qwen2LM, TransformerLM +from paddlespeech.t2s.modules.transformer.activation import Swish +from paddlespeech.t2s.modules.transformer.attention import RelPositionMultiHeadedAttention +from paddlespeech.t2s.modules.transformer.embedding import EspnetRelPositionalEncoding +from paddlespeech.t2s.modules.transformer.subsampling import LinearNoSubsampling + + +COSYVOICE_ACTIVATION_CLASSES = { + "swish": Swish +} +COSYVOICE_SUBSAMPLE_CLASSES = { + "linear": LinearNoSubsampling, +} +COSYVOICE_EMB_CLASSES = { + "rel_pos_espnet": EspnetRelPositionalEncoding, +} +COSYVOICE_ATTENTION_CLASSES = { + "rel_selfattn": RelPositionMultiHeadedAttention, +} + + +# def get_model_type(configs): +# if ( +# isinstance(configs["llm"], TransformerLM) +# and isinstance(configs["flow"], MaskedDiffWithXvec) +# and isinstance(configs["hift"], HiFTGenerator) +# ): +# return CosyVoiceModel +# if ( +# isinstance(configs["llm"], Qwen2LM) +# and isinstance(configs["flow"], CausalMaskedDiffWithXvec) +# and isinstance(configs["hift"], HiFTGenerator) +# ): +# return CosyVoice2Model +# raise TypeError("No valid model type found!") diff --git a/paddlespeech/t2s/models/CosyVoice/common.py b/paddlespeech/t2s/models/CosyVoice/common.py new file mode 100644 index 000000000..9c5a247d9 --- /dev/null +++ b/paddlespeech/t2s/models/CosyVoice/common.py @@ -0,0 +1,216 @@ +import paddle + +"""Unility functions for Transformer.""" +import queue +import random +from typing import List + +import numpy as np + +############################## 相关utils函数,如下 ############################## + +def device2str(type=None, index=None, *, device=None): + type = device if device else type + if isinstance(type, int): + type = f'gpu:{type}' + elif isinstance(type, str): + if 'cuda' in type: + type = type.replace('cuda', 'gpu') + if 'cpu' in type: + type = 'cpu' + elif index is not None: + type = f'{type}:{index}' + elif isinstance(type, paddle.CPUPlace) or (type is None): + type = 'cpu' + elif isinstance(type, paddle.CUDAPlace): + type = f'gpu:{type.get_device_id()}' + + return type +############################## 相关utils函数,如上 ############################## + + +IGNORE_ID = -1 + + +def pad_list(xs: List[paddle.Tensor], pad_value: int): + """Perform padding for the list of tensors. + + Args: + xs (List): List of Tensors [(T_1, `*`), (T_2, `*`), ..., (T_B, `*`)]. + pad_value (float): Value for padding. + + Returns: + Tensor: Padded tensor (B, Tmax, `*`). + + Examples: + >>> x = [torch.ones(4), torch.ones(2), torch.ones(1)] + >>> x + [tensor([1., 1., 1., 1.]), tensor([1., 1.]), tensor([1.])] + >>> pad_list(x, 0) + tensor([[1., 1., 1., 1.], + [1., 1., 0., 0.], + [1., 0., 0., 0.]]) + + """ + max_len = max([len(item) for item in xs]) + batchs = len(xs) + ndim = xs[0].ndim + if ndim == 1: + pad_res = paddle.zeros(batchs, max_len, dtype=xs[0].dtype, device=xs[0].place) + elif ndim == 2: + pad_res = paddle.zeros( + batchs, max_len, xs[0].shape[1], dtype=xs[0].dtype, device=xs[0].place + ) + elif ndim == 3: + pad_res = paddle.zeros( + batchs, + max_len, + xs[0].shape[1], + xs[0].shape[2], + dtype=xs[0].dtype, + device=xs[0].place, + ) + else: + raise ValueError(f"Unsupported ndim: {ndim}") + pad_res.fill_(pad_value) + for i in range(batchs): + pad_res[i, : len(xs[i])] = xs[i] + return pad_res + + +def th_accuracy( + pad_outputs: paddle.Tensor, pad_targets: paddle.Tensor, ignore_label: int +) -> paddle.Tensor: + """Calculate accuracy. + + Args: + pad_outputs (Tensor): Prediction tensors (B * Lmax, D). + pad_targets (LongTensor): Target label tensors (B, Lmax). + ignore_label (int): Ignore label id. + + Returns: + torch.Tensor: Accuracy value (0.0 - 1.0). + + """ + pad_pred = pad_outputs.view( + pad_targets.size(0), pad_targets.size(1), pad_outputs.size(1) + ).argmax(2) + mask = pad_targets != ignore_label + numerator = paddle.sum( + pad_pred.masked_select(mask) == pad_targets.masked_select(mask) + ) + denominator = paddle.sum(mask) + return (numerator / denominator).detach() + + +def get_padding(kernel_size, dilation=1): + return int((kernel_size * dilation - dilation) / 2) + + +def init_weights(m, mean=0.0, std=0.01): + classname = m.__class__.__name__ + if classname.find("Conv") != -1: + m.weight.data.normal_(mean, std) + + +def ras_sampling( + weighted_scores, + decoded_tokens, + sampling, + top_p=0.8, + top_k=25, + win_size=10, + tau_r=0.1, +): + top_ids = nucleus_sampling(weighted_scores, top_p=top_p, top_k=top_k) + rep_num = ( + (paddle.to_tensor(decoded_tokens[-win_size:],dtype = paddle.long).to(weighted_scores.place) == top_ids) + .sum() + .item() + ) + print("top_ids:",top_ids) + if rep_num >= win_size * tau_r: + top_ids = random_sampling(weighted_scores, decoded_tokens, sampling)[0] + return top_ids + + +def nucleus_sampling(weighted_scores, top_p=0.8, top_k=25): + prob, indices = [], [] + cum_prob = 0.0 + sorted_value, sorted_idx = paddle.sort( + descending=True, stable=True, x=weighted_scores.softmax(axis=0) + ), paddle.argsort(descending=True, stable=True, x=weighted_scores.softmax(axis=0)) + + for i in range(len(sorted_idx)): + if cum_prob < top_p and len(prob) < top_k: + cum_prob += sorted_value[i] + prob.append(sorted_value[i]) + indices.append(sorted_idx[i]) + else: + break + prob = paddle.to_tensor(prob).cuda() + indices = paddle.to_tensor(indices, dtype=paddle.long).to(weighted_scores.place) + print("indices:",indices) + # top_ids = indices[prob.multinomial(num_samples=1, replacement=True)] + top_ids = indices[0] + return top_ids + + +def random_sampling(weighted_scores, decoded_tokens, sampling): + top_ids = weighted_scores.softmax(axis=0).multinomial( + num_samples=1, replacement=True + ) + print("random_sampling:",top_ids) + return top_ids + + +def fade_in_out(fade_in_mel, fade_out_mel, window): + device = fade_in_mel.place + fade_in_mel, fade_out_mel = fade_in_mel.cpu(), fade_out_mel.cpu() + mel_overlap_len = int(window.shape[0] / 2) + if fade_in_mel.place == device2str("cpu"): + fade_in_mel = fade_in_mel.clone() + fade_in_mel[..., :mel_overlap_len] = ( + fade_in_mel[..., :mel_overlap_len] * window[:mel_overlap_len] + + fade_out_mel[..., -mel_overlap_len:] * window[mel_overlap_len:] + ) + return fade_in_mel.to(device) + + +def set_all_random_seed(seed): + random.seed(seed) + np.random.seed(seed) + paddle.seed(seed) + paddle.seed(seed) + + +def mask_to_bias(mask: paddle.Tensor, dtype: paddle.dtype) -> paddle.Tensor: + assert mask.dtype == paddle.bool + assert dtype in [paddle.float32, paddle.bfloat16, paddle.float16] + mask = mask.to(dtype) + mask = (1.0 - mask) * -10000000000.0 + return mask + + +class TrtContextWrapper: + def __init__(self, trt_engine, trt_concurrent=1, device="cuda:0"): + self.trt_context_pool = queue.Queue(maxsize=trt_concurrent) + self.trt_engine = trt_engine + for _ in range(trt_concurrent): + trt_context = trt_engine.create_execution_context() + trt_stream = paddle.device.stream_guard( + paddle.device.Stream(device=device2str(device)) + ) + assert ( + trt_context is not None + ), "failed to create trt context, maybe not enough CUDA memory, try reduce current trt concurrent {}".format( + trt_concurrent + ) + self.trt_context_pool.put([trt_context, trt_stream]) + assert self.trt_context_pool.empty() is False, "no avaialbe estimator context" + + def acquire_estimator(self): + return self.trt_context_pool.get(), self.trt_engine + + def release_estimator(self, context, stream): + self.trt_context_pool.put([context, stream]) \ No newline at end of file diff --git a/paddlespeech/t2s/models/CosyVoice/llm.py b/paddlespeech/t2s/models/CosyVoice/llm.py index 6c509d9ab..556d5a0ed 100644 --- a/paddlespeech/t2s/models/CosyVoice/llm.py +++ b/paddlespeech/t2s/models/CosyVoice/llm.py @@ -220,7 +220,7 @@ class TransformerLM(paddle.nn.Layer): num_trials, max_trials = 0, 100 while True: top_ids = self.sampling(weighted_scores, decoded_tokens, sampling) - if not ignore_eos or self.speech_token_size not in top_ids: + if (not ignore_eos) or (top_ids < self.speech_token_size): break num_trials += 1 if num_trials > max_trials: @@ -325,15 +325,18 @@ class Qwen2Encoder(paddle.nn.Layer): ) return outs.hidden_states[-1], masks.unsqueeze(1) - def forward_one_step(self, xs, masks, cache=None): + def forward_one_step(self, xs, masks, cache=None,idx = 0): + input_masks = masks[:, -1, :] outs = self.model( inputs_embeds=xs, attention_mask=input_masks, output_hidden_states=True, return_dict=True, + output_attentions=False, use_cache=True, past_key_values=cache, + index =idx ) xs = outs.hidden_states[-1] new_cache = outs.past_key_values @@ -572,6 +575,7 @@ class Qwen2LM(TransformerLM): out_tokens = [] cache = None for i in range(max_len): + y_pred, cache = self.llm.forward_one_step( lm_input, masks=paddle.tril( @@ -580,6 +584,7 @@ class Qwen2LM(TransformerLM): ) ).to(paddle.bool), cache=cache, + idx = i ) logp = F.log_softmax(self.llm_decoder(y_pred[:, -1]), axis = -1) top_ids = self.sampling_ids( @@ -587,15 +592,13 @@ class Qwen2LM(TransformerLM): out_tokens, sampling, ignore_eos=True if i < min_len else False, - ).item() - if top_ids == self.speech_token_size: + ) + if top_ids in self.stop_token_ids: break - if top_ids > self.speech_token_size: - continue yield top_ids out_tokens.append(top_ids) lm_input = self.speech_embedding.weight[top_ids].reshape([1, 1, -1]) - + print(len(out_tokens)) @paddle.no_grad() def inference_bistream( self, @@ -735,3 +738,5 @@ class Qwen2LM(TransformerLM): raise ValueError("should not get token {}".format(top_ids)) yield top_ids lm_input = self.speech_embedding.weight[top_ids].reshape(1, 1, -1) + + diff --git a/paddlespeech/t2s/models/CosyVoice/mask.py b/paddlespeech/t2s/models/CosyVoice/mask.py new file mode 100644 index 000000000..a27ab4729 --- /dev/null +++ b/paddlespeech/t2s/models/CosyVoice/mask.py @@ -0,0 +1,289 @@ +import paddle + +# ############################## 相关utils函数,如下 ############################## + +# def device2str(type=None, index=None, *, device=None): +# type = device if device else type +# if isinstance(type, int): +# type = f'gpu:{type}' +# elif isinstance(type, str): +# if 'cuda' in type: +# type = type.replace('cuda', 'gpu') +# if 'cpu' in type: +# type = 'cpu' +# elif index is not None: +# type = f'{type}:{index}' +# elif isinstance(type, paddle.CPUPlace) or (type is None): +# type = 'cpu' +# elif isinstance(type, paddle.CUDAPlace): +# type = f'gpu:{type.get_device_id()}' + +# return type + +# def _Tensor_max(self, *args, **kwargs): +# if "other" in kwargs: +# kwargs["y"] = kwargs.pop("other") +# ret = paddle.maximum(self, *args, **kwargs) +# elif len(args) == 1 and isinstance(args[0], paddle.Tensor): +# ret = paddle.maximum(self, *args, **kwargs) +# else: +# if "dim" in kwargs: +# kwargs["axis"] = kwargs.pop("dim") + +# if "axis" in kwargs or len(args) >= 1: +# ret = paddle.max(self, *args, **kwargs), paddle.argmax(self, *args, **kwargs) +# else: +# ret = paddle.max(self, *args, **kwargs) + +# return ret + +# setattr(paddle.Tensor, "_max", _Tensor_max) +# ############################## 相关utils函数,如上 ############################## + + +# """ +# def subsequent_mask( +# size: int, +# device: torch.device = torch.device("cpu"), +# ) -> torch.Tensor: +# ""\"Create mask for subsequent steps (size, size). + +# This mask is used only in decoder which works in an auto-regressive mode. +# This means the current step could only do attention with its left steps. + +# In encoder, fully attention is used when streaming is not necessary and +# the sequence is not long. In this case, no attention mask is needed. + +# When streaming is need, chunk-based attention is used in encoder. See +# subsequent_chunk_mask for the chunk-based attention mask. + +# Args: +# size (int): size of mask +# str device (str): "cpu" or "cuda" or torch.Tensor.device +# dtype (torch.device): result dtype + +# Returns: +# torch.Tensor: mask + +# Examples: +# >>> subsequent_mask(3) +# [[1, 0, 0], +# [1, 1, 0], +# [1, 1, 1]] +# ""\" +# ret = torch.ones(size, size, device=device, dtype=torch.bool) +# return torch.tril(ret) +# """ + + +# def subsequent_mask( +# size: int, device: paddle.device +# ) -> paddle.Tensor: +# """Create mask for subsequent steps (size, size). + +# This mask is used only in decoder which works in an auto-regressive mode. +# This means the current step could only do attention with its left steps. + +# In encoder, fully attention is used when streaming is not necessary and +# the sequence is not long. In this case, no attention mask is needed. + +# When streaming is need, chunk-based attention is used in encoder. See +# subsequent_chunk_mask for the chunk-based attention mask. + +# Args: +# size (int): size of mask +# str device (str): "cpu" or "cuda" or torch.Tensor.device +# dtype (torch.device): result dtype + +# Returns: +# torch.Tensor: mask + +# Examples: +# >>> subsequent_mask(3) +# [[1, 0, 0], +# [1, 1, 0], +# [1, 1, 1]] +# """ +# arange = paddle.arange(size, device=device) +# mask = arange.expand(size, size) +# arange = arange.unsqueeze(-1) +# mask = mask <= arange +# return mask + + +# def subsequent_chunk_mask_deprecated( +# size: int, +# chunk_size: int, +# num_left_chunks: int = -1, +# >>>>>> device: torch.device = device2str("cpu"), +# ) -> paddle.Tensor: +# """Create mask for subsequent steps (size, size) with chunk size, +# this is for streaming encoder + +# Args: +# size (int): size of mask +# chunk_size (int): size of chunk +# num_left_chunks (int): number of left chunks +# <0: use full chunk +# >=0: use num_left_chunks +# device (torch.device): "cpu" or "cuda" or torch.Tensor.device + +# Returns: +# torch.Tensor: mask + +# Examples: +# >>> subsequent_chunk_mask(4, 2) +# [[1, 1, 0, 0], +# [1, 1, 0, 0], +# [1, 1, 1, 1], +# [1, 1, 1, 1]] +# """ +# ret = paddle.zeros(size, size, device=device, dtype=paddle.bool) +# for i in range(size): +# if num_left_chunks < 0: +# start = 0 +# else: +# start = max((i // chunk_size - num_left_chunks) * chunk_size, 0) +# ending = min((i // chunk_size + 1) * chunk_size, size) +# ret[i, start:ending] = True +# return ret + + +# def subsequent_chunk_mask( +# size: int, +# chunk_size: int, +# num_left_chunks: int = -1, +# >>>>>> device: torch.device = device2str("cpu"), +# ) -> paddle.Tensor: +# """Create mask for subsequent steps (size, size) with chunk size, +# this is for streaming encoder + +# Args: +# size (int): size of mask +# chunk_size (int): size of chunk +# num_left_chunks (int): number of left chunks +# <0: use full chunk +# >=0: use num_left_chunks +# device (torch.device): "cpu" or "cuda" or torch.Tensor.device + +# Returns: +# torch.Tensor: mask + +# Examples: +# >>> subsequent_chunk_mask(4, 2) +# [[1, 1, 0, 0], +# [1, 1, 0, 0], +# [1, 1, 1, 1], +# [1, 1, 1, 1]] +# """ +# pos_idx = paddle.arange(size, device=device) +# block_value = ( +# paddle.div(pos_idx, chunk_size, rounding_mode="trunc") + 1 +# ) * chunk_size +# ret = pos_idx.unsqueeze(0) < block_value.unsqueeze(1) +# return ret + + +def add_optional_chunk_mask( + xs: paddle.Tensor, + masks: paddle.Tensor, + use_dynamic_chunk: bool, + use_dynamic_left_chunk: bool, + decoding_chunk_size: int, + static_chunk_size: int, + num_decoding_left_chunks: int, + enable_full_context: bool = True, +): + """Apply optional mask for encoder. + + Args: + xs (torch.Tensor): padded input, (B, L, D), L for max length + mask (torch.Tensor): mask for xs, (B, 1, L) + use_dynamic_chunk (bool): whether to use dynamic chunk or not + use_dynamic_left_chunk (bool): whether to use dynamic left chunk for + training. + decoding_chunk_size (int): decoding chunk size for dynamic chunk, it's + 0: default for training, use random dynamic chunk. + <0: for decoding, use full chunk. + >0: for decoding, use fixed chunk size as set. + static_chunk_size (int): chunk size for static chunk training/decoding + if it's greater than 0, if use_dynamic_chunk is true, + this parameter will be ignored + num_decoding_left_chunks: number of left chunks, this is for decoding, + the chunk size is decoding_chunk_size. + >=0: use num_decoding_left_chunks + <0: use all left chunks + enable_full_context (bool): + True: chunk size is either [1, 25] or full context(max_len) + False: chunk size ~ U[1, 25] + + Returns: + torch.Tensor: chunk mask of the input xs. + """ + if use_dynamic_chunk: + max_len = xs.size(1) + if decoding_chunk_size < 0: + chunk_size = max_len + num_left_chunks = -1 + elif decoding_chunk_size > 0: + chunk_size = decoding_chunk_size + num_left_chunks = num_decoding_left_chunks + else: + chunk_size = paddle.randint(low=1, high=max_len, shape=(1,)).item() + num_left_chunks = -1 + if chunk_size > max_len // 2 and enable_full_context: + chunk_size = max_len + else: + chunk_size = chunk_size % 25 + 1 + if use_dynamic_left_chunk: + max_left_chunks = (max_len - 1) // chunk_size + num_left_chunks = paddle.randint( + low=0, high=max_left_chunks, shape=(1,) + ).item() + chunk_masks = subsequent_chunk_mask( + xs.size(1), chunk_size, num_left_chunks, xs.place + ) + chunk_masks = chunk_masks.unsqueeze(0) + chunk_masks = masks & chunk_masks + elif static_chunk_size > 0: + num_left_chunks = num_decoding_left_chunks + chunk_masks = subsequent_chunk_mask( + xs.size(1), static_chunk_size, num_left_chunks, xs.place + ) + chunk_masks = chunk_masks.unsqueeze(0) + chunk_masks = masks & chunk_masks + else: + chunk_masks = masks + assert chunk_masks.dtype == paddle.bool + if (chunk_masks.sum(dim=-1) == 0).sum().item() != 0: + print( + "get chunk_masks all false at some timestep, force set to true, make sure they are masked in futuer computation!" + ) + chunk_masks[chunk_masks.sum(dim=-1) == 0] = True + return chunk_masks + + +def make_pad_mask(lengths: paddle.Tensor, max_len: int = 0) -> paddle.Tensor: + """Make mask tensor containing indices of padded part. + + See description of make_non_pad_mask. + + Args: + lengths (torch.Tensor): Batch of lengths (B,). + Returns: + torch.Tensor: Mask tensor containing indices of padded part. + + Examples: + >>> lengths = [5, 3, 2] + >>> make_pad_mask(lengths) + masks = [[0, 0, 0, 0 ,0], + [0, 0, 0, 1, 1], + [0, 0, 1, 1, 1]] + """ + batch_size = lengths.shape[0] + max_len = max_len if max_len > 0 else lengths.max().item() + seq_range = paddle.arange(0, max_len, dtype=paddle.int32) + seq_range_expand = seq_range.unsqueeze(0).expand([batch_size, max_len]) + seq_length_expand = lengths.unsqueeze(-1) + mask = seq_range_expand >= seq_length_expand + return mask \ No newline at end of file diff --git a/paddlespeech/t2s/models/__init__.py b/paddlespeech/t2s/models/__init__.py index d8df4368a..c01e79370 100644 --- a/paddlespeech/t2s/models/__init__.py +++ b/paddlespeech/t2s/models/__init__.py @@ -22,3 +22,4 @@ from .transformer_tts import * from .vits import * from .waveflow import * from .wavernn import * +from .CosyVoice import * diff --git a/paddlespeech/t2s/models/hifigan/__init__.py b/paddlespeech/t2s/models/hifigan/__init__.py index 7aa5e9d78..51c0924dd 100644 --- a/paddlespeech/t2s/models/hifigan/__init__.py +++ b/paddlespeech/t2s/models/hifigan/__init__.py @@ -13,3 +13,5 @@ # limitations under the License. from .hifigan import * from .hifigan_updater import * +from .cosy_hifigan import * +from .f0_predictor import * \ No newline at end of file diff --git a/paddlespeech/t2s/models/hifigan/cosy_hifigan.py b/paddlespeech/t2s/models/hifigan/cosy_hifigan.py new file mode 100644 index 000000000..90471c61e --- /dev/null +++ b/paddlespeech/t2s/models/hifigan/cosy_hifigan.py @@ -0,0 +1,443 @@ +import paddle + +"""HIFI-GAN""" +from typing import Dict, List, Optional + +import numpy as np +from scipy.signal import get_window +from paddlespeech.t2s.modules.transformer.activation import Snake +from paddlespeech.t2s.models.CosyVoice.common import get_padding, init_weights + +"""hifigan based generator implementation. + +This code is modified from https://github.com/jik876/hifi-gan + ,https://github.com/kan-bayashi/ParallelWaveGAN and + https://github.com/NVIDIA/BigVGAN + +""" + + +class ResBlock(paddle.nn.Layer): + """Residual block module in HiFiGAN/BigVGAN.""" + + def __init__(self, channels: int=512, kernel_size: int=3, dilations: + List[int]=[1, 3, 5]): + super(ResBlock, self).__init__() + self.convs1 = paddle.nn.LayerList() + self.convs2 = paddle.nn.LayerList() + for dilation in dilations: + self.convs1.append(paddle.nn.utils.weight_norm(layer=paddle.nn. + Conv1D(channels, channels, kernel_size, 1, dilation= + dilation, padding=get_padding(kernel_size, dilation)))) + self.convs2.append(paddle.nn.utils.weight_norm(layer=paddle.nn. + Conv1D(channels, channels, kernel_size, 1, dilation=1, + padding=get_padding(kernel_size, 1)))) + self.convs1.apply(init_weights) + self.convs2.apply(init_weights) + self.activations1 = paddle.nn.LayerList(sublayers=[Snake(channels, + alpha_logscale=False) for _ in range(len(self.convs1))]) + self.activations2 = paddle.nn.LayerList(sublayers=[Snake(channels, + alpha_logscale=False) for _ in range(len(self.convs2))]) + + def forward(self, x: paddle.Tensor) ->paddle.Tensor: + for idx in range(len(self.convs1)): + xt = self.activations1[idx](x) + xt = self.convs1[idx](xt) + xt = self.activations2[idx](xt) + xt = self.convs2[idx](xt) + x = xt + x + return x + + def remove_weight_norm(self): + for idx in range(len(self.convs1)): + paddle.nn.utils.remove_weight_norm(layer=self.convs1[idx]) + paddle.nn.utils.remove_weight_norm(layer=self.convs2[idx]) + + +class SineGen(paddle.nn.Layer): + """ Definition of sine generator + SineGen(samp_rate, harmonic_num = 0, + sine_amp = 0.1, noise_std = 0.003, + voiced_threshold = 0, + flag_for_pulse=False) + samp_rate: sampling rate in Hz + harmonic_num: number of harmonic overtones (default 0) + sine_amp: amplitude of sine-wavefrom (default 0.1) + noise_std: std of Gaussian noise (default 0.003) + voiced_thoreshold: F0 threshold for U/V classification (default 0) + flag_for_pulse: this SinGen is used inside PulseGen (default False) + Note: when flag_for_pulse is True, the first time step of a voiced + segment is always sin(np.pi) or cos(0) + """ + + def __init__(self, samp_rate, harmonic_num=0, sine_amp=0.1, noise_std= + 0.003, voiced_threshold=0): + super(SineGen, self).__init__() + self.sine_amp = sine_amp + self.noise_std = noise_std + self.harmonic_num = harmonic_num + self.sampling_rate = samp_rate + self.voiced_threshold = voiced_threshold + + def _f02uv(self, f0): + uv = (f0 > self.voiced_threshold).astype(paddle.float32) + return uv + + @paddle.no_grad() + def forward(self, f0): + """ + :param f0: [B, 1, sample_len], Hz + :return: [B, 1, sample_len] + """ + F_mat = paddle.zeros([f0.size(0), self.harmonic_num + 1, f0.size(-1)]).to(f0.place) + for i in range(self.harmonic_num + 1): + F_mat[:, i:i + 1, :] = f0 * (i + 1) / self.sampling_rate + theta_mat = 2 * np.pi * (paddle.cumsum(F_mat, axis=-1) % 1) + u_dist = paddle.distribution.Uniform(low=-np.pi, high=np.pi) + phase_vec = u_dist.sample(shape=(f0.size(0), self.harmonic_num + 1, 1) + ).to(F_mat.place) + phase_vec[:, 0, :] = 0 + sine_waves = self.sine_amp * paddle.sin(theta_mat + phase_vec) + uv = self._f02uv(f0) + noise_amp = uv * self.noise_std + (1 - uv) * self.sine_amp / 3 + noise = noise_amp * paddle.randn(shape=sine_waves.shape, dtype= + sine_waves.dtype) + sine_waves = sine_waves * uv + noise + return sine_waves, uv, noise + + +class SourceModuleHnNSF(paddle.nn.Layer): + """ SourceModule for hn-nsf + SourceModule(sampling_rate, harmonic_num=0, sine_amp=0.1, + add_noise_std=0.003, voiced_threshod=0) + sampling_rate: sampling_rate in Hz + harmonic_num: number of harmonic above F0 (default: 0) + sine_amp: amplitude of sine source signal (default: 0.1) + add_noise_std: std of additive Gaussian noise (default: 0.003) + note that amplitude of noise in unvoiced is decided + by sine_amp + voiced_threshold: threhold to set U/V given F0 (default: 0) + Sine_source, noise_source = SourceModuleHnNSF(F0_sampled) + F0_sampled (batchsize, length, 1) + Sine_source (batchsize, length, 1) + noise_source (batchsize, length 1) + uv (batchsize, length, 1) + """ + + def __init__(self, sampling_rate, upsample_scale, harmonic_num=0, + sine_amp=0.1, add_noise_std=0.003, voiced_threshod=0): + super(SourceModuleHnNSF, self).__init__() + self.sine_amp = sine_amp + self.noise_std = add_noise_std + self.l_sin_gen = SineGen(sampling_rate, harmonic_num, sine_amp, + add_noise_std, voiced_threshod) + self.l_linear = paddle.nn.Linear(in_features=harmonic_num + 1, + out_features=1) + self.l_tanh = paddle.nn.Tanh() + + def forward(self, x): + """ + Sine_source, noise_source = SourceModuleHnNSF(F0_sampled) + F0_sampled (batchsize, length, 1) + Sine_source (batchsize, length, 1) + noise_source (batchsize, length 1) + """ + with paddle.no_grad(): + sine_wavs, uv, _ = self.l_sin_gen(paddle.transpose(x,perm=[0,2,1])) + sine_wavs = paddle.transpose(sine_wavs,perm=[0,2,1]) + uv = paddle.transpose(uv,perm=[0,2,1]) + sine_merge = self.l_tanh(self.l_linear(sine_wavs)) + + noise = paddle.randn(shape=uv.shape, dtype=uv.dtype + ) * self.sine_amp / 3 + return sine_merge, noise, uv + + +class SineGen2(paddle.nn.Layer): + """ Definition of sine generator + SineGen(samp_rate, harmonic_num = 0, + sine_amp = 0.1, noise_std = 0.003, + voiced_threshold = 0, + flag_for_pulse=False) + samp_rate: sampling rate in Hz + harmonic_num: number of harmonic overtones (default 0) + sine_amp: amplitude of sine-wavefrom (default 0.1) + noise_std: std of Gaussian noise (default 0.003) + voiced_thoreshold: F0 threshold for U/V classification (default 0) + flag_for_pulse: this SinGen is used inside PulseGen (default False) + Note: when flag_for_pulse is True, the first time step of a voiced + segment is always sin(np.pi) or cos(0) + """ + + def __init__(self, samp_rate, upsample_scale, harmonic_num=0, sine_amp= + 0.1, noise_std=0.003, voiced_threshold=0, flag_for_pulse=False): + super(SineGen2, self).__init__() + self.sine_amp = sine_amp + self.noise_std = noise_std + self.harmonic_num = harmonic_num + self.axis = self.harmonic_num + 1 + self.sampling_rate = samp_rate + self.voiced_threshold = voiced_threshold + self.flag_for_pulse = flag_for_pulse + self.upsample_scale = upsample_scale + + def _f02uv(self, f0): + uv = (f0 > self.voiced_threshold).astype(paddle.float32) + return uv + + def _f02sine(self, f0_values): + """ f0_values: (batchsize, length, axis) + where axis indicates fundamental tone and overtones + """ + rad_values = f0_values / self.sampling_rate % 1 + rand_ini = paddle.rand(shape=[f0_values.shape[0], f0_values.shape[2]]) + rand_ini[:, 0] = 0 + rad_values[:, 0, :] = rad_values[:, 0, :] + rand_ini + if not self.flag_for_pulse: + x = paddle.transpose(rad_values,perm = [0,2,1]) + + rad_values = paddle.transpose(paddle.nn.functional.interpolate(x=x, scale_factor=1 / self.upsample_scale, mode='linear'),perm = [0,2,1]) + phase = paddle.cumsum(rad_values, axis=1) * 2 * np.pi + phase = paddle.transpose(paddle.nn.functional.interpolate(x=paddle.transpose(phase,perm = [0,2,1]) * self.upsample_scale, scale_factor=int(self.upsample_scale),mode='linear'),perm = [0,2,1]) + sines = paddle.sin(phase) + else: + uv = self._f02uv(f0_values) + uv_1 = paddle.roll(uv, shifts=-1, axis=1) + uv_1[:, -1, :] = 1 + u_loc = (uv < 1) * (uv_1 > 0) + tmp_cumsum = paddle.cumsum(rad_values, axis=1) + for idx in range(f0_values.shape[0]): + temp_sum = tmp_cumsum[idx, u_loc[idx, :, 0], :] + temp_sum[1:, :] = temp_sum[1:, :] - temp_sum[0:-1, :] + tmp_cumsum[idx, :, :] = 0 + tmp_cumsum[idx, u_loc[idx, :, 0], :] = temp_sum + i_phase = paddle.cumsum(rad_values - tmp_cumsum, axis=1) + sines = paddle.cos(i_phase * 2 * np.pi) + return sines + + def forward(self, f0): + """ sine_tensor, uv = forward(f0) + input F0: tensor(batchsize=1, length, axis=1) + f0 for unvoiced steps should be 0 + output sine_tensor: tensor(batchsize=1, length, axis) + output uv: tensor(batchsize=1, length, 1) + """ + paddle.seed(1986) + fn = paddle.multiply(f0, paddle.to_tensor([[range(1, self.harmonic_num + + 2)]],dtype='float32',place=f0.place)) + + sine_waves = self._f02sine(fn) * self.sine_amp + uv = self._f02uv(f0) + noise_amp = uv * self.noise_std + (1 - uv) * self.sine_amp / 3 + noise = noise_amp * paddle.randn(shape=sine_waves.shape, dtype= + sine_waves.dtype) + sine_waves = sine_waves * uv + noise + + return sine_waves, uv, noise + + +class SourceModuleHnNSF2(paddle.nn.Layer): + """ SourceModule for hn-nsf + SourceModule(sampling_rate, harmonic_num=0, sine_amp=0.1, + add_noise_std=0.003, voiced_threshod=0) + sampling_rate: sampling_rate in Hz + harmonic_num: number of harmonic above F0 (default: 0) + sine_amp: amplitude of sine source signal (default: 0.1) + add_noise_std: std of additive Gaussian noise (default: 0.003) + note that amplitude of noise in unvoiced is decided + by sine_amp + voiced_threshold: threhold to set U/V given F0 (default: 0) + Sine_source, noise_source = SourceModuleHnNSF(F0_sampled) + F0_sampled (batchsize, length, 1) + Sine_source (batchsize, length, 1) + noise_source (batchsize, length 1) + uv (batchsize, length, 1) + """ + + def __init__(self, sampling_rate, upsample_scale, harmonic_num=0, + sine_amp=0.1, add_noise_std=0.003, voiced_threshod=0): + super(SourceModuleHnNSF2, self).__init__() + self.sine_amp = sine_amp + self.noise_std = add_noise_std + self.l_sin_gen = SineGen2(sampling_rate, upsample_scale, + harmonic_num, sine_amp, add_noise_std, voiced_threshod) + self.l_linear = paddle.nn.Linear(in_features=harmonic_num + 1, + out_features=1) + self.l_tanh = paddle.nn.Tanh() + + def forward(self, x): + """ + Sine_source, noise_source = SourceModuleHnNSF(F0_sampled) + F0_sampled (batchsize, length, 1) + Sine_source (batchsize, length, 1) + noise_source (batchsize, length 1) + """ + paddle.seed(1986) + with paddle.no_grad(): + sine_wavs, uv, _ = self.l_sin_gen(x) + sine_merge = self.l_tanh(self.l_linear(sine_wavs)) + noise = paddle.randn(shape=uv.shape, dtype=uv.dtype + ) * self.sine_amp / 3 + return sine_merge, noise, uv + + +class HiFTGenerator(paddle.nn.Layer): + """ + HiFTNet Generator: Neural Source Filter + ISTFTNet + https://arxiv.org/abs/2309.09493 + """ + + def __init__(self, in_channels: int=80, base_channels: int=512, + nb_harmonics: int=8, sampling_rate: int=22050, nsf_alpha: float=0.1, + nsf_sigma: float=0.003, nsf_voiced_threshold: float=10, + upsample_rates: List[int]=[8, 8], upsample_kernel_sizes: List[int]= + [16, 16], istft_params: Dict[str, int]={'n_fft': 16, 'hop_len': 4}, + resblock_kernel_sizes: List[int]=[3, 7, 11], + resblock_dilation_sizes: List[List[int]]=[[1, 3, 5], [1, 3, 5], [1, + 3, 5]], source_resblock_kernel_sizes: List[int]=[7, 11], + source_resblock_dilation_sizes: List[List[int]]=[[1, 3, 5], [1, 3, + 5]], lrelu_slope: float=0.1, audio_limit: float=0.99, f0_predictor: + paddle.nn.Layer=None): + super(HiFTGenerator, self).__init__() + self.out_channels = 1 + self.nb_harmonics = nb_harmonics + self.sampling_rate = sampling_rate + self.istft_params = istft_params + self.lrelu_slope = lrelu_slope + self.audio_limit = audio_limit + self.num_kernels = len(resblock_kernel_sizes) + self.num_upsamples = len(upsample_rates) + this_SourceModuleHnNSF = (SourceModuleHnNSF if self.sampling_rate == + 22050 else SourceModuleHnNSF2) + self.m_source = this_SourceModuleHnNSF(sampling_rate=sampling_rate, + upsample_scale=np.prod(upsample_rates) * istft_params['hop_len' + ], harmonic_num=nb_harmonics, sine_amp=nsf_alpha, add_noise_std + =nsf_sigma, voiced_threshod=nsf_voiced_threshold) + self.f0_upsamp = paddle.nn.Upsample(scale_factor=(1,int(np.prod( upsample_rates) * istft_params['hop_len']))) + self.conv_pre = paddle.nn.utils.weight_norm(layer=paddle.nn.Conv1D(in_channels = in_channels, out_channels = base_channels, kernel_size=7, stride=1, padding=3)) + self.ups = paddle.nn.LayerList() + for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)): + self.ups.append(paddle.nn.utils.weight_norm(layer=paddle.nn. + Conv1DTranspose(in_channels=base_channels // 2 ** i, + out_channels=base_channels // 2 ** (i + 1), kernel_size=k, + stride=u, padding=(k - u) // 2))) + self.source_downs = paddle.nn.LayerList() + self.source_resblocks = paddle.nn.LayerList() + downsample_rates = [1] + upsample_rates[::-1][:-1] + downsample_cum_rates = np.cumprod(downsample_rates) + for i, (u, k, d) in enumerate(zip(downsample_cum_rates[::-1], + source_resblock_kernel_sizes, source_resblock_dilation_sizes)): + if u == 1: + self.source_downs.append(paddle.nn.Conv1D(istft_params[ + 'n_fft'] + 2, base_channels // 2 ** (i + 1), 1, 1)) + else: + self.source_downs.append(paddle.nn.Conv1D( + in_channels=istft_params['n_fft'] + 2, + out_channels=base_channels // (2 ** (i + 1)), + kernel_size=(u * 2,), + stride=(u,), + padding=int(u // 2), + )) + self.source_resblocks.append(ResBlock(base_channels // 2 ** (i + + 1), k, d)) + self.resblocks = paddle.nn.LayerList() + for i in range(len(self.ups)): + ch = base_channels // 2 ** (i + 1) + for _, (k, d) in enumerate(zip(resblock_kernel_sizes, + resblock_dilation_sizes)): + self.resblocks.append(ResBlock(ch, k, d)) + self.conv_post = paddle.nn.utils.weight_norm(layer=paddle.nn.Conv1D + (ch, istft_params['n_fft'] + 2, 7, 1, padding=3)) + self.ups.apply(init_weights) + self.conv_post.apply(init_weights) + self.reflection_pad = paddle.nn.Pad1D(padding=(1, 0), mode='reflect') + self.stft_window = paddle.to_tensor(get_window('hann', + istft_params['n_fft'], fftbins=True).astype(np.float32)) + self.f0_predictor = f0_predictor + + def remove_weight_norm(self): + print('Removing weight norm...') + for l in self.ups: + paddle.nn.utils.remove_weight_norm(layer=l) + for l in self.resblocks: + l.remove_weight_norm() + paddle.nn.utils.remove_weight_norm(layer=self.conv_pre) + paddle.nn.utils.remove_weight_norm(layer=self.conv_post) + self.m_source.remove_weight_norm() + for l in self.source_downs: + paddle.nn.utils.remove_weight_norm(layer=l) + for l in self.source_resblocks: + l.remove_weight_norm() + + def _stft(self, x): + spec = paddle.signal.stft(x=x, n_fft=self.istft_params['n_fft'], + hop_length=self.istft_params['hop_len'], win_length=self. + istft_params['n_fft'], window=self.stft_window.to(x.place)) + spec = paddle.as_real(spec) + return spec[..., 0], spec[..., 1] + + def _istft(self, magnitude, phase): + magnitude = paddle.clip(magnitude, max=100.0) + real = magnitude * paddle.cos(phase) + img = magnitude * paddle.sin(phase) + inverse_transform = paddle.signal.istft(x=paddle.complex(real, img), + n_fft=self.istft_params['n_fft'], hop_length=self.istft_params[ + 'hop_len'], win_length=self.istft_params['n_fft'], window=self. + stft_window.to(magnitude.place)) + return inverse_transform + + def decode(self, x: paddle.Tensor, s: paddle.Tensor=paddle.zeros([1, 1, 0]) + ) ->paddle.Tensor: + s_stft_real, s_stft_imag = self._stft(s.squeeze(1)) + s_stft = paddle.cat([s_stft_real, s_stft_imag], dim=1) + x = self.conv_pre(x) + for i in range(self.num_upsamples): + x = paddle.nn.functional.leaky_relu(x=x, negative_slope=self. + lrelu_slope) + x = self.ups[i](x) + if i == self.num_upsamples - 1: + x = self.reflection_pad(x) + si = self.source_downs[i](s_stft) + si = self.source_resblocks[i](si) + x = x + si + xs = None + for j in range(self.num_kernels): + if xs is None: + xs = self.resblocks[i * self.num_kernels + j](x) + else: + xs += self.resblocks[i * self.num_kernels + j](x) + x = xs / self.num_kernels + x = paddle.nn.functional.leaky_relu(x=x) + x = self.conv_post(x) + magnitude = paddle.exp(x=x[:, :self.istft_params['n_fft'] // 2 + 1, :]) + phase = paddle.sin(x[:, self.istft_params['n_fft'] // 2 + 1:, :]) + x = self._istft(magnitude, phase) + x = paddle.clip(x, -self.audio_limit, self.audio_limit) + return x + + def forward(self, batch: dict) ->Dict[str, + Optional[paddle.Tensor]]: + speech_feat = paddle.transpose(batch['speech_feat'],perm = [0,2,1]).to(device) + f0 = self.f0_predictor(speech_feat) + s = paddle.transpose(self.f0_upsamp(f0[:, None]),perm = [0,2,1]) + s, _, _ = self.m_source(s) + s = paddle.transpose(s,perm = [0,2,1]) + + generated_speech = self.decode(x=speech_feat, s=s) + return generated_speech, f0 + + @paddle.no_grad() + def inference(self, speech_feat: paddle.Tensor, cache_source: paddle. + Tensor=paddle.zeros([1, 1, 0])) ->paddle.Tensor: + paddle.seed(1986) + f0 = self.f0_predictor(speech_feat) + f0_4d = f0[:, None].unsqueeze(2) + s_4d = self.f0_upsamp(f0_4d) + s_3d = s_4d.squeeze(2) + s = paddle.transpose(s_3d, perm=[0, 2, 1]) + s, _, _ = self.m_source(s) + s = paddle.transpose(s,perm = [0,2,1]) + if cache_source.shape[2] != 0: + s[:, :, :cache_source.shape[2]] = cache_source + generated_speech = self.decode(x=speech_feat, s=s) + return generated_speech, s diff --git a/paddlespeech/t2s/models/hifigan/f0_predictor.py b/paddlespeech/t2s/models/hifigan/f0_predictor.py new file mode 100644 index 000000000..8cf686181 --- /dev/null +++ b/paddlespeech/t2s/models/hifigan/f0_predictor.py @@ -0,0 +1,25 @@ +import paddle +class ConvRNNF0Predictor(paddle.nn.Layer): + + def __init__(self, num_class: int=1, in_channels: int=80, cond_channels: + int=512): + super().__init__() + self.num_class = num_class + self.condnet = paddle.nn.Sequential(paddle.nn.utils.weight_norm( + layer=paddle.nn.Conv1D(in_channels, cond_channels, kernel_size= + 3, padding=1)), paddle.nn.ELU(), paddle.nn.utils.weight_norm( + layer=paddle.nn.Conv1D(cond_channels, cond_channels, + kernel_size=3, padding=1)), paddle.nn.ELU(), paddle.nn.utils. + weight_norm(layer=paddle.nn.Conv1D(cond_channels, cond_channels, + kernel_size=3, padding=1)), paddle.nn.ELU(), paddle.nn.utils. + weight_norm(layer=paddle.nn.Conv1D(cond_channels, cond_channels, + kernel_size=3, padding=1)), paddle.nn.ELU(), paddle.nn.utils. + weight_norm(layer=paddle.nn.Conv1D(cond_channels, cond_channels, + kernel_size=3, padding=1)), paddle.nn.ELU()) + self.classifier = paddle.nn.Linear(in_features=cond_channels, + out_features=self.num_class) + + def forward(self, x: paddle.Tensor) ->paddle.Tensor: + x = self.condnet(x) + x = paddle.transpose(x, perm=[0, 2, 1]) + return paddle.abs(x=self.classifier(x).squeeze(-1)) diff --git a/paddlespeech/t2s/modules/conv.py b/paddlespeech/t2s/modules/conv.py index 922af03f2..484dd10f7 100644 --- a/paddlespeech/t2s/modules/conv.py +++ b/paddlespeech/t2s/modules/conv.py @@ -1,4 +1,4 @@ -# Copyright (c) 2020 PaddlePaddle Authors. All Rights Reserved. + # Copyright (c) 2020 PaddlePaddle Authors. All Rights Reserved. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. @@ -124,7 +124,7 @@ class Conv1dCell(nn.Conv1D): self._reshaped_weight = paddle.reshape(self.weight, (self._out_channels, -1)) - def initialize_buffer(self, x_t): + def initialize_buffer(self, x_t): """Initialize the buffer for the step input. Args: diff --git a/paddlespeech/t2s/modules/flow/__init__.py b/paddlespeech/t2s/modules/flow/__init__.py new file mode 100644 index 000000000..2e3d46a07 --- /dev/null +++ b/paddlespeech/t2s/modules/flow/__init__.py @@ -0,0 +1 @@ +from .flow import CausalMaskedDiffWithXvec \ No newline at end of file diff --git a/paddlespeech/t2s/modules/flow/attention.py b/paddlespeech/t2s/modules/flow/attention.py index b5a7069d3..f4f87c4fc 100644 --- a/paddlespeech/t2s/modules/flow/attention.py +++ b/paddlespeech/t2s/modules/flow/attention.py @@ -1,6 +1,93 @@ +from typing import Any, Dict, Optional +import paddle +from .attention_processor import Attention -class BasicTransformerBlock(nn.Module): - r""" +def _chunked_feed_forward( + ff: paddle.nn.Layer, hidden_states: paddle.Tensor, chunk_dim: int, chunk_size: int +): + if hidden_states.shape[chunk_dim] % chunk_size != 0: + raise ValueError( + f"`hidden_states` dimension to be chunked: {hidden_states.shape[chunk_dim]} has to be divisible by chunk size: {chunk_size}. Make sure to set an appropriate `chunk_size` when calling `unet.enable_forward_chunking`." + ) + num_chunks = hidden_states.shape[chunk_dim] // chunk_size + ff_output = paddle.cat( + [ff(hid_slice) for hid_slice in hidden_states.chunk(num_chunks, dim=chunk_dim)], + dim=chunk_dim, + ) + return ff_output + +class SinusoidalPositionalEmbedding(paddle.nn.Layer): + """Apply positional information to a sequence of embeddings. + + Takes in a sequence of embeddings with shape (batch_size, seq_length, embed_dim) and adds positional embeddings to + them + + Args: + embed_dim: (int): Dimension of the positional embedding. + max_seq_length: Maximum sequence length to apply positional embeddings + + """ + + def __init__(self, embed_dim: int, max_seq_length: int = 32): + super().__init__() + position = paddle.arange(max_seq_length).unsqueeze(1) + div_term = paddle.exp( + x=paddle.arange(0, embed_dim, 2) * (-math.log(10000.0) / embed_dim) + ) + pe = paddle.zeros(1, max_seq_length, embed_dim) + pe[0, :, 0::2] = paddle.sin(position * div_term) + pe[0, :, 1::2] = paddle.cos(position * div_term) + self.register_buffer(name="pe", tensor=pe) + + def forward(self, x): + _, seq_length, _ = x.shape + x = x + self.pe[:, :seq_length] + return x + +class GatedSelfAttentionDense(paddle.nn.Layer): + """ + A gated self-attention dense layer that combines visual features and object features. + + Parameters: + query_dim (`int`): The number of channels in the query. + context_dim (`int`): The number of channels in the context. + n_heads (`int`): The number of heads to use for attention. + d_head (`int`): The number of channels in each head. + """ + + def __init__(self, query_dim: int, context_dim: int, n_heads: int, d_head: int): + super().__init__() + self.linear = paddle.nn.Linear(in_features=context_dim, out_features=query_dim) + self.attn = Attention(query_dim=query_dim, heads=n_heads, dim_head=d_head) + self.ff = FeedForward(query_dim, activation_fn="geglu") + self.norm1 = paddle.nn.LayerNorm(normalized_shape=query_dim) + self.norm2 = paddle.nn.LayerNorm(normalized_shape=query_dim) + self.add_parameter( + name="alpha_attn", + parameter=paddle.nn.parameter.Parameter(paddle.tensor(0.0)), + ) + self.add_parameter( + name="alpha_dense", + parameter=paddle.nn.parameter.Parameter(paddle.tensor(0.0)), + ) + self.enabled = True + + def forward(self, x: paddle.Tensor, objs: paddle.Tensor) -> paddle.Tensor: + if not self.enabled: + return x + n_visual = x.shape[1] + objs = self.linear(objs) + x = ( + x + + self.alpha_attn.tanh() + * self.attn(self.norm1(paddle.cat([x, objs], dim=1)))[:, :n_visual, :] + ) + x = x + self.alpha_dense.tanh() * self.ff(self.norm2(x)) + return x + + +class BasicTransformerBlock(paddle.nn.Layer): + """ A basic Transformer block. Parameters: @@ -9,15 +96,29 @@ class BasicTransformerBlock(nn.Module): attention_head_dim (`int`): The number of channels in each head. dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use. cross_attention_dim (`int`, *optional*): The size of the encoder_hidden_states vector for cross attention. - only_cross_attention (`bool`, *optional*): - Whether to use only cross-attention layers. In this case two cross attention layers are used. - double_self_attention (`bool`, *optional*): - Whether to use two self-attention layers. In this case no cross attention layers are used. activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward. num_embeds_ada_norm (: obj: `int`, *optional*): The number of diffusion steps used during training. See `Transformer2DModel`. attention_bias (: obj: `bool`, *optional*, defaults to `False`): Configure if the attentions should contain a bias parameter. + only_cross_attention (`bool`, *optional*): + Whether to use only cross-attention layers. In this case two cross attention layers are used. + double_self_attention (`bool`, *optional*): + Whether to use two self-attention layers. In this case no cross attention layers are used. + upcast_attention (`bool`, *optional*): + Whether to upcast the attention computation to float32. This is useful for mixed precision training. + norm_elementwise_affine (`bool`, *optional*, defaults to `True`): + Whether to use learnable elementwise affine parameters for normalization. + norm_type (`str`, *optional*, defaults to `"layer_norm"`): + The normalization layer to use. Can be `"layer_norm"`, `"ada_norm"` or `"ada_norm_zero"`. + final_dropout (`bool` *optional*, defaults to False): + Whether to apply a final dropout after the last feed-forward layer. + attention_type (`str`, *optional*, defaults to `"default"`): + The type of attention to use. Can be `"default"` or `"gated"` or `"gated-text-image"`. + positional_embeddings (`str`, *optional*, defaults to `None`): + The type of positional embeddings to apply to. + num_positional_embeddings (`int`, *optional*, defaults to `None`): + The maximum number of positional embeddings to apply. """ def __init__( @@ -26,202 +127,260 @@ class BasicTransformerBlock(nn.Module): num_attention_heads: int, attention_head_dim: int, dropout=0.0, - activation_fn: str = "geglu", cross_attention_dim: Optional[int] = None, + activation_fn: str = "geglu", num_embeds_ada_norm: Optional[int] = None, attention_bias: bool = False, + only_cross_attention: bool = False, double_self_attention: bool = False, upcast_attention: bool = False, norm_elementwise_affine: bool = True, norm_type: str = "layer_norm", + norm_eps: float = 1e-05, final_dropout: bool = False, + attention_type: str = "default", + positional_embeddings: Optional[str] = None, + num_positional_embeddings: Optional[int] = None, + ada_norm_continous_conditioning_embedding_dim: Optional[int] = None, + ada_norm_bias: Optional[int] = None, + ff_inner_dim: Optional[int] = None, + ff_bias: bool = True, + attention_out_bias: bool = True, ): super().__init__() - self.use_ada_layer_norm_zero = (num_embeds_ada_norm is not None) and norm_type == "ada_norm_zero" - self.use_ada_layer_norm = (num_embeds_ada_norm is not None) and norm_type == "ada_norm" - + self.only_cross_attention = only_cross_attention + self.use_ada_layer_norm_zero = ( + num_embeds_ada_norm is not None and norm_type == "ada_norm_zero" + ) + self.use_ada_layer_norm = ( + num_embeds_ada_norm is not None and norm_type == "ada_norm" + ) + self.use_ada_layer_norm_single = norm_type == "ada_norm_single" + self.use_layer_norm = norm_type == "layer_norm" + self.use_ada_layer_norm_continuous = norm_type == "ada_norm_continuous" if norm_type in ("ada_norm", "ada_norm_zero") and num_embeds_ada_norm is None: raise ValueError( - f"`norm_type` is set to {norm_type}, but `num_embeds_ada_norm` is not defined. Please make sure to" - f" define `num_embeds_ada_norm` if setting `norm_type` to {norm_type}." + f"`norm_type` is set to {norm_type}, but `num_embeds_ada_norm` is not defined. Please make sure to define `num_embeds_ada_norm` if setting `norm_type` to {norm_type}." + ) + self.norm_type = norm_type + self.num_embeds_ada_norm = num_embeds_ada_norm + if positional_embeddings and num_positional_embeddings is None: + raise ValueError( + "If `positional_embedding` type is defined, `num_positition_embeddings` must also be defined." ) - # Define 3 blocks. Each block has its own normalization layer. - # 1. Self-Attn - self.norm1 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine) + if positional_embeddings == "sinusoidal": + self.pos_embed = SinusoidalPositionalEmbedding( + dim, max_seq_length=num_positional_embeddings + ) + else: + self.pos_embed = None + + + self.norm1 = paddle.nn.LayerNorm( + normalized_shape=dim, + weight_attr=norm_elementwise_affine, + bias_attr=norm_elementwise_affine, + epsilon=norm_eps, + ) self.attn1 = Attention( query_dim=dim, heads=num_attention_heads, dim_head=attention_head_dim, dropout=dropout, bias=attention_bias, - cross_attention_dim=None - upcast_attention=False + cross_attention_dim=cross_attention_dim if only_cross_attention else None, + upcast_attention=upcast_attention, + out_bias=attention_out_bias, ) - # 2. Cross-Attn - self.norm2 = None - self.attn2 = None - self.norm3 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine) - self.ff = FeedForward(dim, dropout=dropout, activation_fn=activation_fn, final_dropout=final_dropout) - - # let chunk size default to None + if cross_attention_dim is not None or double_self_attention: + if norm_type == "ada_norm": + self.norm2 = AdaLayerNorm(dim, num_embeds_ada_norm) + elif norm_type == "ada_norm_continuous": + self.norm2 = AdaLayerNormContinuous( + dim, + ada_norm_continous_conditioning_embedding_dim, + norm_elementwise_affine, + norm_eps, + ada_norm_bias, + "rms_norm", + ) + else: + self.norm2 = paddle.nn.LayerNorm( + normalized_shape=dim, + epsilon=norm_eps, + weight_attr=norm_elementwise_affine, + bias_attr=norm_elementwise_affine, + ) + self.attn2 = Attention( + query_dim=dim, + cross_attention_dim=cross_attention_dim + if not double_self_attention + else None, + heads=num_attention_heads, + dim_head=attention_head_dim, + dropout=dropout, + bias=attention_bias, + upcast_attention=upcast_attention, + out_bias=attention_out_bias, + ) + else: + self.norm2 = None + self.attn2 = None + if norm_type == "ada_norm_continuous": + self.norm3 = AdaLayerNormContinuous( + dim, + ada_norm_continous_conditioning_embedding_dim, + norm_elementwise_affine, + norm_eps, + ada_norm_bias, + "layer_norm", + ) + elif norm_type in [ + "ada_norm_zero", + "ada_norm", + "layer_norm", + "ada_norm_continuous", + ]: + self.norm3 = paddle.nn.LayerNorm( + normalized_shape=dim, + epsilon=norm_eps, + weight_attr=norm_elementwise_affine, + bias_attr=norm_elementwise_affine, + ) + elif norm_type == "layer_norm_i2vgen": + self.norm3 = None + self.ff = FeedForward( + dim, + dropout=dropout, + activation_fn=activation_fn, + final_dropout=final_dropout, + inner_dim=ff_inner_dim, + bias=ff_bias, + ) + if attention_type == "gated" or attention_type == "gated-text-image": + self.fuser = GatedSelfAttentionDense( + dim, cross_attention_dim, num_attention_heads, attention_head_dim + ) + if norm_type == "ada_norm_single": + self.scale_shift_table = paddle.nn.parameter.Parameter( + paddle.randn(6, dim) / dim**0.5 + ) self._chunk_size = None self._chunk_dim = 0 - def forward(self,hidden_states): - norm_hidden_states = self.norm1(hidden_states) - cross_attention_kwargs = {} + + def set_chunk_feed_forward(self, chunk_size: Optional[int], dim: int = 0): + self._chunk_size = chunk_size + self._chunk_dim = dim + + def forward( + self, + hidden_states: paddle.Tensor, + attention_mask: Optional[paddle.Tensor] = None, + encoder_hidden_states: Optional[paddle.Tensor] = None, + encoder_attention_mask: Optional[paddle.Tensor] = None, + timestep: Optional[paddle.Tensor] = None, + cross_attention_kwargs: Dict[str, Any] = None, + class_labels: Optional[paddle.Tensor] = None, + added_cond_kwargs: Optional[Dict[str, paddle.Tensor]] = None, + ) -> paddle.Tensor: + if cross_attention_kwargs is not None: + if cross_attention_kwargs.get("scale", None) is not None: + logger.warning( + "Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored." + ) + batch_size = hidden_states.shape[0] + if self.norm_type == "ada_norm": + norm_hidden_states = self.norm1(hidden_states, timestep) + elif self.norm_type == "ada_norm_zero": + (norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp) = self.norm1( + hidden_states, timestep, class_labels, hidden_dtype=hidden_states.dtype + ) + elif self.norm_type in ["layer_norm", "layer_norm_i2vgen"]: + norm_hidden_states = self.norm1(hidden_states) + elif self.norm_type == "ada_norm_continuous": + norm_hidden_states = self.norm1( + hidden_states, added_cond_kwargs["pooled_text_emb"] + ) + elif self.norm_type == "ada_norm_single": + shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = ( + self.scale_shift_table[None] + timestep.reshape(batch_size, 6, -1) + ).chunk(6, dim=1) + norm_hidden_states = self.norm1(hidden_states) + norm_hidden_states = norm_hidden_states * (1 + scale_msa) + shift_msa + norm_hidden_states = norm_hidden_states.squeeze(1) + else: + raise ValueError("Incorrect norm used") + if self.pos_embed is not None: + norm_hidden_states = self.pos_embed(norm_hidden_states) + cross_attention_kwargs = ( + cross_attention_kwargs.copy() if cross_attention_kwargs is not None else {} + ) + gligen_kwargs = cross_attention_kwargs.pop("gligen", None) attn_output = self.attn1( norm_hidden_states, - encoder_hidden_states=encoder_hidden_states if self.only_cross_attention else None, - attention_mask=encoder_attention_mask if self.only_cross_attention else attention_mask, + encoder_hidden_states=encoder_hidden_states + if self.only_cross_attention + else None, + attention_mask=attention_mask, **cross_attention_kwargs, ) + if self.norm_type == "ada_norm_zero": + attn_output = gate_msa.unsqueeze(1) * attn_output + elif self.norm_type == "ada_norm_single": + attn_output = gate_msa * attn_output hidden_states = attn_output + hidden_states - norm_hidden_states = self.norm3(hidden_states) - ff_output = self.ff(norm_hidden_states) - hidden_states = ff_output + hidden_states - return hidden_states - -class FeedForward(nn.Layer): - - def __init__( - self, - dim: int, - dim_out: Optional[int] = None, - mult: int = 4, - dropout: float = 0.0, - activation_fn: str = "geglu", - final_dropout: bool = False, - ): - super().__init__() - inner_dim = int(dim * mult) - dim_out = dim_out if dim_out is not None else dim - - if activation_fn == "gelu": - act_fn = GELU(dim, inner_dim, approximate=False) - elif activation_fn == "gelu-approximate": - act_fn = GELU(dim, inner_dim, approximate=True) - elif activation_fn == "geglu": - act_fn = GEGLU(dim, inner_dim) - elif activation_fn == "geglu-approximate": - act_fn = ApproximateGELU(dim, inner_dim) - elif activation_fn == "snakebeta": - act_fn = SnakeBeta(dim, inner_dim) + if hidden_states.ndim == 4: + hidden_states = hidden_states.squeeze(1) + if gligen_kwargs is not None: + hidden_states = self.fuser(hidden_states, gligen_kwargs["objs"]) + if self.attn2 is not None: + if self.norm_type == "ada_norm": + norm_hidden_states = self.norm2(hidden_states, timestep) + elif self.norm_type in ["ada_norm_zero", "layer_norm", "layer_norm_i2vgen"]: + norm_hidden_states = self.norm2(hidden_states) + elif self.norm_type == "ada_norm_single": + norm_hidden_states = hidden_states + elif self.norm_type == "ada_norm_continuous": + norm_hidden_states = self.norm2( + hidden_states, added_cond_kwargs["pooled_text_emb"] + ) + else: + raise ValueError("Incorrect norm") + if self.pos_embed is not None and self.norm_type != "ada_norm_single": + norm_hidden_states = self.pos_embed(norm_hidden_states) + attn_output = self.attn2( + norm_hidden_states, + encoder_hidden_states=encoder_hidden_states, + attention_mask=encoder_attention_mask, + **cross_attention_kwargs, + ) + hidden_states = attn_output + hidden_states + if self.norm_type == "ada_norm_continuous": + norm_hidden_states = self.norm3( + hidden_states, added_cond_kwargs["pooled_text_emb"] + ) + elif not self.norm_type == "ada_norm_single": + norm_hidden_states = self.norm3(hidden_states) + if self.norm_type == "ada_norm_zero": + norm_hidden_states = ( + norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] + ) + if self.norm_type == "ada_norm_single": + norm_hidden_states = self.norm2(hidden_states) + norm_hidden_states = norm_hidden_states * (1 + scale_mlp) + shift_mlp + if self._chunk_size is not None: + ff_output = _chunked_feed_forward( + self.ff, norm_hidden_states, self._chunk_dim, self._chunk_size + ) else: - act_fn = GEGLU(dim, inner_dim) - - self.net = nn.LayerList() - self.net.append(act_fn) - self.net.append(nn.Dropout(dropout)) - self.net.append(LoRACompatibleLinear(inner_dim, dim_out)) - - if final_dropout: - self.net.append(nn.Dropout(dropout)) - - def forward(self, hidden_states): - for module in self.net: - hidden_states = module(hidden_states) + ff_output = self.ff(norm_hidden_states) + if self.norm_type == "ada_norm_zero": + ff_output = gate_mlp.unsqueeze(1) * ff_output + elif self.norm_type == "ada_norm_single": + ff_output = gate_mlp * ff_output + hidden_states = ff_output + hidden_states + if hidden_states.ndim == 4: + hidden_states = hidden_states.squeeze(1) return hidden_states -query_dim=dim, -heads=num_attention_heads, -dim_head=attention_head_dim, -dropout=dropout, -bias=attention_bias, -cross_attention_dim=None, -upcast_attention=upcast_attention, -class Attention(nn.Module): - def __init__( - self, - query_dim: int, - cross_attention_dim: Optional[int] = None, - heads: int = 8, - dim_head: int = 64, - dropout: float = 0.0, - bias: bool = False, - upcast_attention: bool = False, - upcast_softmax: bool = False, - cross_attention_norm: Optional[str] = None, - cross_attention_norm_num_groups: int = 32, - qk_norm: Optional[str] = None, - norm_num_groups: Optional[int] = None, - spatial_norm_dim: Optional[int] = None, - out_bias: bool = True, - scale_qk: bool = True, - only_cross_attention: bool = False, - eps: float = 1e-5, - rescale_output_factor: float = 1.0, - processor: Optional["AttnProcessor"] = None, - out_dim: int = None, - ): - super().__init__() - self.inner_dim = out_dim if out_dim is not None else dim_head * heads - self.query_dim = query_dim - self.use_bias = bias - self.is_cross_attention = cross_attention_dim is not None - self.cross_attention_dim = cross_attention_dim if cross_attention_dim is not None else query_dim - self.upcast_attention = upcast_attention - self.upcast_softmax = upcast_softmax - self.rescale_output_factor = rescale_output_factor - self.dropout = dropout - self.fused_projections = False - self.out_dim = out_dim if out_dim is not None else query_dim - # we make use of this private variable to know whether this class is loaded - # with an deprecated state dict so that we can convert it on the fly - - self.scale_qk = scale_qk - self.scale = dim_head**-0.5 if self.scale_qk else 1.0 - - self.heads = out_dim // dim_head if out_dim is not None else heads - # for slice_size > 0 the attention score computation - # is split across the batch axis to save memory - # You can set slice_size with `set_attention_slice` - self.sliceable_head_dim = heads - self.to_q = nn.Linear(query_dim, self.inner_dim, bias=bias) - self.to_k = nn.Linear(self.cross_attention_dim, self.inner_dim, bias=bias) - self.to_v = nn.Linear(self.cross_attention_dim, self.inner_dim, bias=bias) - self.to_out = nn.ModuleList([]) - self.to_out.append(nn.Linear(self.inner_dim, self.out_dim, bias=out_bias)) - self.to_out.append(nn.Dropout(dropout)) - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: Optional[torch.Tensor] = None, - attention_mask: Optional[torch.Tensor] = None, - temb: Optional[torch.Tensor] = None, - ): - residual = hidden_states - input_ndim = hidden_states.ndim - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - query = attn.to_q(hidden_states) - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states = hidden_states.to(query.dtype) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if input_ndim == 4: - hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) - hidden_states = hidden_states / attn.rescale_output_factor - return hidden_states diff --git a/paddlespeech/t2s/modules/flow/attention_processor.py b/paddlespeech/t2s/modules/flow/attention_processor.py new file mode 100644 index 000000000..f34998083 --- /dev/null +++ b/paddlespeech/t2s/modules/flow/attention_processor.py @@ -0,0 +1,625 @@ +import inspect +import math +from typing import Callable, List, Optional, Union +import paddle + +class Attention(paddle.nn.Layer): + """ + A cross attention layer. + + Parameters: + query_dim (`int`): + The number of channels in the query. + cross_attention_dim (`int`, *optional*): + The number of channels in the encoder_hidden_states. If not given, defaults to `query_dim`. + heads (`int`, *optional*, defaults to 8): + The number of heads to use for multi-head attention. + dim_head (`int`, *optional*, defaults to 64): + The number of channels in each head. + dropout (`float`, *optional*, defaults to 0.0): + The dropout probability to use. + bias (`bool`, *optional*, defaults to False): + Set to `True` for the query, key, and value linear layers to contain a bias parameter. + upcast_attention (`bool`, *optional*, defaults to False): + Set to `True` to upcast the attention computation to `float32`. + upcast_softmax (`bool`, *optional*, defaults to False): + Set to `True` to upcast the softmax computation to `float32`. + cross_attention_norm (`str`, *optional*, defaults to `None`): + The type of normalization to use for the cross attention. Can be `None`, `layer_norm`, or `group_norm`. + cross_attention_norm_num_groups (`int`, *optional*, defaults to 32): + The number of groups to use for the group norm in the cross attention. + added_kv_proj_dim (`int`, *optional*, defaults to `None`): + The number of channels to use for the added key and value projections. If `None`, no projection is used. + norm_num_groups (`int`, *optional*, defaults to `None`): + The number of groups to use for the group norm in the attention. + spatial_norm_dim (`int`, *optional*, defaults to `None`): + The number of channels to use for the spatial normalization. + out_bias (`bool`, *optional*, defaults to `True`): + Set to `True` to use a bias in the output linear layer. + scale_qk (`bool`, *optional*, defaults to `True`): + Set to `True` to scale the query and key by `1 / sqrt(dim_head)`. + only_cross_attention (`bool`, *optional*, defaults to `False`): + Set to `True` to only use cross attention and not added_kv_proj_dim. Can only be set to `True` if + `added_kv_proj_dim` is not `None`. + eps (`float`, *optional*, defaults to 1e-5): + An additional value added to the denominator in group normalization that is used for numerical stability. + rescale_output_factor (`float`, *optional*, defaults to 1.0): + A factor to rescale the output by dividing it with this value. + residual_connection (`bool`, *optional*, defaults to `False`): + Set to `True` to add the residual connection to the output. + _from_deprecated_attn_block (`bool`, *optional*, defaults to `False`): + Set to `True` if the attention block is loaded from a deprecated state dict. + processor (`AttnProcessor`, *optional*, defaults to `None`): + The attention processor to use. If `None`, defaults to `AttnProcessor2_0` if `torch 2.x` is used and + `AttnProcessor` otherwise. + """ + + def __init__( + self, + query_dim: int, + cross_attention_dim: Optional[int] = None, + heads: int = 8, + dim_head: int = 64, + dropout: float = 0.0, + bias: bool = False, + upcast_attention: bool = False, + upcast_softmax: bool = False, + cross_attention_norm: Optional[str] = None, + cross_attention_norm_num_groups: int = 32, + qk_norm: Optional[str] = None, + added_kv_proj_dim: Optional[int] = None, + norm_num_groups: Optional[int] = None, + spatial_norm_dim: Optional[int] = None, + out_bias: bool = True, + scale_qk: bool = True, + only_cross_attention: bool = False, + eps: float = 1e-05, + rescale_output_factor: float = 1.0, + residual_connection: bool = False, + _from_deprecated_attn_block: bool = False, + processor: Optional["AttnProcessor"] = None, + out_dim: int = None, + context_pre_only=None, + ): + super().__init__() + self.inner_dim = out_dim if out_dim is not None else dim_head * heads + self.query_dim = query_dim + self.use_bias = bias + self.is_cross_attention = cross_attention_dim is not None + self.cross_attention_dim = ( + cross_attention_dim if cross_attention_dim is not None else query_dim + ) + self.upcast_attention = upcast_attention + self.upcast_softmax = upcast_softmax + self.rescale_output_factor = rescale_output_factor + self.residual_connection = residual_connection + self.dropout = dropout + self.fused_projections = False + self.out_dim = out_dim if out_dim is not None else query_dim + self.context_pre_only = context_pre_only + self._from_deprecated_attn_block = _from_deprecated_attn_block + self.scale_qk = scale_qk + self.scale = dim_head**-0.5 if self.scale_qk else 1.0 + self.heads = out_dim // dim_head if out_dim is not None else heads + self.sliceable_head_dim = heads + self.added_kv_proj_dim = added_kv_proj_dim + self.only_cross_attention = only_cross_attention + if self.added_kv_proj_dim is None and self.only_cross_attention: + raise ValueError( + "`only_cross_attention` can only be set to True if `added_kv_proj_dim` is not None. Make sure to set either `only_cross_attention=False` or define `added_kv_proj_dim`." + ) + if norm_num_groups is not None: + self.group_norm = paddle.nn.GroupNorm( + num_channels=query_dim, + num_groups=norm_num_groups, + epsilon=eps, + weight_attr=True, + bias_attr=True, + ) + else: + self.group_norm = None + if spatial_norm_dim is not None: + self.spatial_norm = SpatialNorm( + f_channels=query_dim, zq_channels=spatial_norm_dim + ) + else: + self.spatial_norm = None + if qk_norm is None: + self.norm_q = None + self.norm_k = None + elif qk_norm == "layer_norm": + self.norm_q = paddle.nn.LayerNorm(normalized_shape=dim_head, epsilon=eps) + self.norm_k = paddle.nn.LayerNorm(normalized_shape=dim_head, epsilon=eps) + else: + raise ValueError( + f"unknown qk_norm: {qk_norm}. Should be None or 'layer_norm'" + ) + if cross_attention_norm is None: + self.norm_cross = None + elif cross_attention_norm == "layer_norm": + self.norm_cross = paddle.nn.LayerNorm( + normalized_shape=self.cross_attention_dim + ) + elif cross_attention_norm == "group_norm": + if self.added_kv_proj_dim is not None: + norm_cross_num_channels = added_kv_proj_dim + else: + norm_cross_num_channels = self.cross_attention_dim + self.norm_cross = paddle.nn.GroupNorm( + num_channels=norm_cross_num_channels, + num_groups=cross_attention_norm_num_groups, + epsilon=1e-05, + weight_attr=True, + bias_attr=True, + ) + else: + raise ValueError( + f"unknown cross_attention_norm: {cross_attention_norm}. Should be None, 'layer_norm' or 'group_norm'" + ) + self.to_q = paddle.nn.Linear( + in_features=query_dim, out_features=self.inner_dim, bias_attr=bias + ) + if not self.only_cross_attention: + self.to_k = paddle.nn.Linear( + in_features=self.cross_attention_dim, + out_features=self.inner_dim, + bias_attr=bias, + ) + self.to_v = paddle.nn.Linear( + in_features=self.cross_attention_dim, + out_features=self.inner_dim, + bias_attr=bias, + ) + else: + self.to_k = None + self.to_v = None + if self.added_kv_proj_dim is not None: + self.add_k_proj = paddle.nn.Linear( + in_features=added_kv_proj_dim, out_features=self.inner_dim + ) + self.add_v_proj = paddle.nn.Linear( + in_features=added_kv_proj_dim, out_features=self.inner_dim + ) + if self.context_pre_only is not None: + self.add_q_proj = paddle.nn.Linear( + in_features=added_kv_proj_dim, out_features=self.inner_dim + ) + self.to_out = paddle.nn.LayerList(sublayers=[]) + self.to_out.append( + paddle.nn.Linear( + in_features=self.inner_dim, + out_features=self.out_dim, + bias_attr=out_bias, + ) + ) + self.to_out.append(paddle.nn.Dropout(p=dropout)) + if self.context_pre_only is not None and not self.context_pre_only: + self.to_add_out = paddle.nn.Linear( + in_features=self.inner_dim, + out_features=self.out_dim, + bias_attr=out_bias, + ) + processor = AttnProcessor() + processor = AttnProcessor2_0() + self.set_processor(processor) + + def set_processor(self, processor: "AttnProcessor") -> None: + """ + Set the attention processor to use. + + Args: + processor (`AttnProcessor`): + The attention processor to use. + """ + if ( + hasattr(self, "processor") + and isinstance(self.processor, paddle.nn.Layer) + and not isinstance(processor, paddle.nn.Layer) + ): + self._modules.pop("processor") + self.processor = processor + + def forward( + self, + hidden_states: paddle.Tensor, + encoder_hidden_states: Optional[paddle.Tensor] = None, + attention_mask: Optional[paddle.Tensor] = None, + **cross_attention_kwargs, + ) -> paddle.Tensor: + """ + The forward method of the `Attention` class. + + Args: + hidden_states (`torch.Tensor`): + The hidden states of the query. + encoder_hidden_states (`torch.Tensor`, *optional*): + The hidden states of the encoder. + attention_mask (`torch.Tensor`, *optional*): + The attention mask to use. If `None`, no mask is applied. + **cross_attention_kwargs: + Additional keyword arguments to pass along to the cross attention. + + Returns: + `torch.Tensor`: The output of the attention layer. + """ + attn_parameters = set( + inspect.signature(self.processor.__call__).parameters.keys() + ) + quiet_attn_parameters = {"ip_adapter_masks"} + unused_kwargs = [ + k + for k, _ in cross_attention_kwargs.items() + if k not in attn_parameters and k not in quiet_attn_parameters + ] + if len(unused_kwargs) > 0: + logger.warning( + f"cross_attention_kwargs {unused_kwargs} are not expected by {self.processor.__class__.__name__} and will be ignored." + ) + cross_attention_kwargs = { + k: w for k, w in cross_attention_kwargs.items() if k in attn_parameters + } + return self.processor( + self, + hidden_states, + encoder_hidden_states=encoder_hidden_states, + attention_mask=attention_mask, + **cross_attention_kwargs, + ) + + def batch_to_head_dim(self, tensor: paddle.Tensor) -> paddle.Tensor: + """ + Reshape the tensor from `[batch_size, seq_len, dim]` to `[batch_size // heads, seq_len, dim * heads]`. `heads` + is the number of heads initialized while constructing the `Attention` class. + + Args: + tensor (`torch.Tensor`): The tensor to reshape. + + Returns: + `torch.Tensor`: The reshaped tensor. + """ + head_size = self.heads + batch_size, seq_len, dim = tensor.shape + tensor = tensor.reshape(batch_size // head_size, head_size, seq_len, dim) + tensor = tensor.permute(0, 2, 1, 3).reshape( + batch_size // head_size, seq_len, dim * head_size + ) + return tensor + + def head_to_batch_dim( + self, tensor: paddle.Tensor, out_dim: int = 3 + ) -> paddle.Tensor: + """ + Reshape the tensor from `[batch_size, seq_len, dim]` to `[batch_size, seq_len, heads, dim // heads]` `heads` is + the number of heads initialized while constructing the `Attention` class. + + Args: + tensor (`torch.Tensor`): The tensor to reshape. + out_dim (`int`, *optional*, defaults to `3`): The output dimension of the tensor. If `3`, the tensor is + reshaped to `[batch_size * heads, seq_len, dim // heads]`. + + Returns: + `torch.Tensor`: The reshaped tensor. + """ + head_size = self.heads + if tensor.ndim == 3: + batch_size, seq_len, dim = tensor.shape + extra_dim = 1 + else: + batch_size, extra_dim, seq_len, dim = tensor.shape + tensor = tensor.reshape( + batch_size, seq_len * extra_dim, head_size, dim // head_size + ) + tensor = tensor.permute(0, 2, 1, 3) + if out_dim == 3: + tensor = tensor.reshape( + batch_size * head_size, seq_len * extra_dim, dim // head_size + ) + return tensor + + def get_attention_scores( + self, + query: paddle.Tensor, + key: paddle.Tensor, + attention_mask: paddle.Tensor = None, + ) -> paddle.Tensor: + """ + Compute the attention scores. + + Args: + query (`torch.Tensor`): The query tensor. + key (`torch.Tensor`): The key tensor. + attention_mask (`torch.Tensor`, *optional*): The attention mask to use. If `None`, no mask is applied. + + Returns: + `torch.Tensor`: The attention probabilities/scores. + """ + dtype = query.dtype + if self.upcast_attention: + query = query.float() + key = key.float() + if attention_mask is None: + baddbmm_input = paddle.empty( + query.shape[0], + query.shape[1], + key.shape[1], + dtype=query.dtype, + device=query.place, + ) + beta = 0 + else: + baddbmm_input = attention_mask + beta = 1 + attention_scores = paddle.baddbmm( + input=baddbmm_input, + x=query, + y=key.transpose(-1, -2), + beta=beta, + alpha=self.scale, + ) + del baddbmm_input + if self.upcast_softmax: + attention_scores = attention_scores.float() + attention_probs = attention_scores.softmax(dim=-1) + del attention_scores + attention_probs = attention_probs.to(dtype) + return attention_probs + + def prepare_attention_mask( + self, + attention_mask: paddle.Tensor, + target_length: int, + batch_size: int, + out_dim: int = 3, + ) -> paddle.Tensor: + """ + Prepare the attention mask for the attention computation. + + Args: + attention_mask (`torch.Tensor`): + The attention mask to prepare. + target_length (`int`): + The target length of the attention mask. This is the length of the attention mask after padding. + batch_size (`int`): + The batch size, which is used to repeat the attention mask. + out_dim (`int`, *optional*, defaults to `3`): + The output dimension of the attention mask. Can be either `3` or `4`. + + Returns: + `torch.Tensor`: The prepared attention mask. + """ + head_size = self.heads + if attention_mask is None: + return attention_mask + current_length: int = attention_mask.shape[-1] + if current_length != target_length: + if attention_mask.device.type == "mps": + padding_shape = ( + attention_mask.shape[0], + attention_mask.shape[1], + target_length, + ) + padding = paddle.zeros( + padding_shape, + dtype=attention_mask.dtype, + device=attention_mask.place, + ) + attention_mask = paddle.cat([attention_mask, padding], dim=2) + else: + attention_mask = paddle.compat.pad( + attention_mask, (0, target_length), value=0.0 + ) + if out_dim == 3: + if attention_mask.shape[0] < batch_size * head_size: + attention_mask = attention_mask.repeat_interleave(head_size, axis=0) + elif out_dim == 4: + attention_mask = attention_mask.unsqueeze(1) + attention_mask = attention_mask.repeat_interleave(head_size, dim=1) + return attention_mask + + def norm_encoder_hidden_states( + self, encoder_hidden_states: paddle.Tensor + ) -> paddle.Tensor: + """ + Normalize the encoder hidden states. Requires `self.norm_cross` to be specified when constructing the + `Attention` class. + + Args: + encoder_hidden_states (`torch.Tensor`): Hidden states of the encoder. + + Returns: + `torch.Tensor`: The normalized encoder hidden states. + """ + assert ( + self.norm_cross is not None + ), "self.norm_cross must be defined to call self.norm_encoder_hidden_states" + if isinstance(self.norm_cross, paddle.nn.LayerNorm): + encoder_hidden_states = self.norm_cross(encoder_hidden_states) + elif isinstance(self.norm_cross, paddle.nn.GroupNorm): + encoder_hidden_states = encoder_hidden_states.transpose(1, 2) + encoder_hidden_states = self.norm_cross(encoder_hidden_states) + encoder_hidden_states = encoder_hidden_states.transpose(1, 2) + else: + assert False + return encoder_hidden_states + + @paddle.no_grad() + def fuse_projections(self, fuse=True): + device = self.to_q.weight.data.place + dtype = self.to_q.weight.data.dtype + if not self.is_cross_attention: + concatenated_weights = paddle.cat( + [self.to_q.weight.data, self.to_k.weight.data, self.to_v.weight.data] + ) + in_features = concatenated_weights.shape[1] + out_features = concatenated_weights.shape[0] + self.to_qkv = paddle.nn.Linear( + in_features=in_features, + out_features=out_features, + bias_attr=self.use_bias, + ) + self.to_qkv.weight.copy_(concatenated_weights) + if self.use_bias: + concatenated_bias = paddle.cat( + [self.to_q.bias.data, self.to_k.bias.data, self.to_v.bias.data] + ) + self.to_qkv.bias.copy_(concatenated_bias) + else: + concatenated_weights = paddle.cat( + [self.to_k.weight.data, self.to_v.weight.data] + ) + in_features = concatenated_weights.shape[1] + out_features = concatenated_weights.shape[0] + self.to_kv = paddle.nn.Linear( + in_features=in_features, + out_features=out_features, + bias_attr=self.use_bias, + ) + self.to_kv.weight.copy_(concatenated_weights) + if self.use_bias: + concatenated_bias = paddle.cat( + [self.to_k.bias.data, self.to_v.bias.data] + ) + self.to_kv.bias.copy_(concatenated_bias) + self.fused_projections = fuse + + +class AttnProcessor: + """ + Default processor for performing attention-related computations. + """ + + def __call__( + self, + attn: Attention, + hidden_states: paddle.Tensor, + encoder_hidden_states: Optional[paddle.Tensor] = None, + attention_mask: Optional[paddle.Tensor] = None, + temb: Optional[paddle.Tensor] = None, + *args, + **kwargs, + ) -> paddle.Tensor: + residual = hidden_states + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + input_ndim = hidden_states.ndim + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view( + batch_size, channel, height * width + ).transpose(1, 2) + batch_size, sequence_length, _ = ( + hidden_states.shape + if encoder_hidden_states is None + else encoder_hidden_states.shape + ) + attention_mask = attn.prepare_attention_mask( + attention_mask, sequence_length, batch_size + ) + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose( + 1, 2 + ) + query = attn.to_q(hidden_states) + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states( + encoder_hidden_states + ) + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + query = attn.head_to_batch_dim(query) + key = attn.head_to_batch_dim(key) + value = attn.head_to_batch_dim(value) + attention_probs = attn.get_attention_scores(query, key, attention_mask) + hidden_states = paddle.bmm(attention_probs, value) + hidden_states = attn.batch_to_head_dim(hidden_states) + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape( + batch_size, channel, height, width + ) + if attn.residual_connection: + hidden_states = hidden_states + residual + hidden_states = hidden_states / attn.rescale_output_factor + return hidden_states + +class AttnProcessor2_0: + """ + Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). + """ + + def __init__(self): + pass + def __call__( + self, + attn: Attention, + hidden_states: paddle.Tensor, + encoder_hidden_states: Optional[paddle.Tensor] = None, + attention_mask: Optional[paddle.Tensor] = None, + temb: Optional[paddle.Tensor] = None, + *args, + **kwargs, + ) -> paddle.Tensor: + residual = hidden_states + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + input_ndim = hidden_states.ndim + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view( + batch_size, channel, height * width + ).transpose(1, 2) + batch_size, sequence_length, _ = ( + hidden_states.shape + if encoder_hidden_states is None + else encoder_hidden_states.shape + ) + if attention_mask is not None: + attention_mask = attn.prepare_attention_mask( + attention_mask, sequence_length, batch_size + ) + attention_mask = attention_mask.view( + [batch_size, attn.heads, -1, attention_mask.shape[-1]] + ) + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose( + 1, 2 + ) + query = attn.to_q(hidden_states) + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states( + encoder_hidden_states + ) + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + inner_dim = key.shape[-1] + head_dim = inner_dim // attn.heads + query = paddle.transpose(query.view([batch_size, -1, attn.heads, head_dim]),perm = [0,2,1]) + key = paddle.transpose(key.view([batch_size, -1, attn.heads, head_dim]),perm = [0,2,1]) + value =paddle.transpose(value.view([batch_size, -1, attn.heads, head_dim]),perm = [0,2,1]) + hidden_states = paddle.nn.functional.scaled_dot_product_attention( + query.transpose([0, 2, 1, 3]), + key.transpose([0, 2, 1, 3]), + value.transpose([0, 2, 1, 3]), + attn_mask=attention_mask, + dropout_p=0.0, + is_causal=False, + ).transpose([0, 2, 1, 3]) + hidden_states = paddle.transpose(hidden_states,perm = [0,2,1]).reshape( + [batch_size, -1, attn.heads * head_dim] + ) + hidden_states = hidden_states.to(query.dtype) + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + if input_ndim == 4: + hidden_states =paddle.transpose(hidden_states,perm = [0,1,3,2]).reshape( + [batch_size, channel, height, width] + ) + if attn.residual_connection: + hidden_states = hidden_states + residual + hidden_states = hidden_states / attn.rescale_output_factor + return hidden_states diff --git a/paddlespeech/t2s/modules/flow/attention_processor_back.py b/paddlespeech/t2s/modules/flow/attention_processor_back.py new file mode 100644 index 000000000..5621efe33 --- /dev/null +++ b/paddlespeech/t2s/modules/flow/attention_processor_back.py @@ -0,0 +1,3015 @@ +import inspect +import math +from importlib import import_module +from typing import Callable, List, Optional, Union + +import paddle + +from ..image_processor import IPAdapterMaskProcessor +from ..utils import deprecate, logging +from ..utils.import_utils import is_torch_npu_available, is_xformers_available +from ..utils.torch_utils import maybe_allow_in_graph +from .lora import LoRALinearLayer + +logger = logging.get_logger(__name__) +if is_torch_npu_available(): + import torch_npu +if is_xformers_available(): + import xformers + import xformers.ops +else: + xformers = None + + +@maybe_allow_in_graph +class Attention(paddle.nn.Layer): + """ + A cross attention layer. + + Parameters: + query_dim (`int`): + The number of channels in the query. + cross_attention_dim (`int`, *optional*): + The number of channels in the encoder_hidden_states. If not given, defaults to `query_dim`. + heads (`int`, *optional*, defaults to 8): + The number of heads to use for multi-head attention. + dim_head (`int`, *optional*, defaults to 64): + The number of channels in each head. + dropout (`float`, *optional*, defaults to 0.0): + The dropout probability to use. + bias (`bool`, *optional*, defaults to False): + Set to `True` for the query, key, and value linear layers to contain a bias parameter. + upcast_attention (`bool`, *optional*, defaults to False): + Set to `True` to upcast the attention computation to `float32`. + upcast_softmax (`bool`, *optional*, defaults to False): + Set to `True` to upcast the softmax computation to `float32`. + cross_attention_norm (`str`, *optional*, defaults to `None`): + The type of normalization to use for the cross attention. Can be `None`, `layer_norm`, or `group_norm`. + cross_attention_norm_num_groups (`int`, *optional*, defaults to 32): + The number of groups to use for the group norm in the cross attention. + added_kv_proj_dim (`int`, *optional*, defaults to `None`): + The number of channels to use for the added key and value projections. If `None`, no projection is used. + norm_num_groups (`int`, *optional*, defaults to `None`): + The number of groups to use for the group norm in the attention. + spatial_norm_dim (`int`, *optional*, defaults to `None`): + The number of channels to use for the spatial normalization. + out_bias (`bool`, *optional*, defaults to `True`): + Set to `True` to use a bias in the output linear layer. + scale_qk (`bool`, *optional*, defaults to `True`): + Set to `True` to scale the query and key by `1 / sqrt(dim_head)`. + only_cross_attention (`bool`, *optional*, defaults to `False`): + Set to `True` to only use cross attention and not added_kv_proj_dim. Can only be set to `True` if + `added_kv_proj_dim` is not `None`. + eps (`float`, *optional*, defaults to 1e-5): + An additional value added to the denominator in group normalization that is used for numerical stability. + rescale_output_factor (`float`, *optional*, defaults to 1.0): + A factor to rescale the output by dividing it with this value. + residual_connection (`bool`, *optional*, defaults to `False`): + Set to `True` to add the residual connection to the output. + _from_deprecated_attn_block (`bool`, *optional*, defaults to `False`): + Set to `True` if the attention block is loaded from a deprecated state dict. + processor (`AttnProcessor`, *optional*, defaults to `None`): + The attention processor to use. If `None`, defaults to `AttnProcessor2_0` if `torch 2.x` is used and + `AttnProcessor` otherwise. + """ + + def __init__( + self, + query_dim: int, + cross_attention_dim: Optional[int] = None, + heads: int = 8, + dim_head: int = 64, + dropout: float = 0.0, + bias: bool = False, + upcast_attention: bool = False, + upcast_softmax: bool = False, + cross_attention_norm: Optional[str] = None, + cross_attention_norm_num_groups: int = 32, + qk_norm: Optional[str] = None, + added_kv_proj_dim: Optional[int] = None, + norm_num_groups: Optional[int] = None, + spatial_norm_dim: Optional[int] = None, + out_bias: bool = True, + scale_qk: bool = True, + only_cross_attention: bool = False, + eps: float = 1e-05, + rescale_output_factor: float = 1.0, + residual_connection: bool = False, + _from_deprecated_attn_block: bool = False, + processor: Optional["AttnProcessor"] = None, + out_dim: int = None, + context_pre_only=None, + ): + super().__init__() + self.inner_dim = out_dim if out_dim is not None else dim_head * heads + self.query_dim = query_dim + self.use_bias = bias + self.is_cross_attention = cross_attention_dim is not None + self.cross_attention_dim = ( + cross_attention_dim if cross_attention_dim is not None else query_dim + ) + self.upcast_attention = upcast_attention + self.upcast_softmax = upcast_softmax + self.rescale_output_factor = rescale_output_factor + self.residual_connection = residual_connection + self.dropout = dropout + self.fused_projections = False + self.out_dim = out_dim if out_dim is not None else query_dim + self.context_pre_only = context_pre_only + self._from_deprecated_attn_block = _from_deprecated_attn_block + self.scale_qk = scale_qk + self.scale = dim_head**-0.5 if self.scale_qk else 1.0 + self.heads = out_dim // dim_head if out_dim is not None else heads + self.sliceable_head_dim = heads + self.added_kv_proj_dim = added_kv_proj_dim + self.only_cross_attention = only_cross_attention + if self.added_kv_proj_dim is None and self.only_cross_attention: + raise ValueError( + "`only_cross_attention` can only be set to True if `added_kv_proj_dim` is not None. Make sure to set either `only_cross_attention=False` or define `added_kv_proj_dim`." + ) + if norm_num_groups is not None: + self.group_norm = paddle.nn.GroupNorm( + num_channels=query_dim, + num_groups=norm_num_groups, + epsilon=eps, + weight_attr=True, + bias_attr=True, + ) + else: + self.group_norm = None + if spatial_norm_dim is not None: + self.spatial_norm = SpatialNorm( + f_channels=query_dim, zq_channels=spatial_norm_dim + ) + else: + self.spatial_norm = None + if qk_norm is None: + self.norm_q = None + self.norm_k = None + elif qk_norm == "layer_norm": + self.norm_q = paddle.nn.LayerNorm(normalized_shape=dim_head, epsilon=eps) + self.norm_k = paddle.nn.LayerNorm(normalized_shape=dim_head, epsilon=eps) + else: + raise ValueError( + f"unknown qk_norm: {qk_norm}. Should be None or 'layer_norm'" + ) + if cross_attention_norm is None: + self.norm_cross = None + elif cross_attention_norm == "layer_norm": + self.norm_cross = paddle.nn.LayerNorm( + normalized_shape=self.cross_attention_dim + ) + elif cross_attention_norm == "group_norm": + if self.added_kv_proj_dim is not None: + norm_cross_num_channels = added_kv_proj_dim + else: + norm_cross_num_channels = self.cross_attention_dim + self.norm_cross = paddle.nn.GroupNorm( + num_channels=norm_cross_num_channels, + num_groups=cross_attention_norm_num_groups, + epsilon=1e-05, + weight_attr=True, + bias_attr=True, + ) + else: + raise ValueError( + f"unknown cross_attention_norm: {cross_attention_norm}. Should be None, 'layer_norm' or 'group_norm'" + ) + self.to_q = paddle.nn.Linear( + in_features=query_dim, out_features=self.inner_dim, bias_attr=bias + ) + if not self.only_cross_attention: + self.to_k = paddle.nn.Linear( + in_features=self.cross_attention_dim, + out_features=self.inner_dim, + bias_attr=bias, + ) + self.to_v = paddle.nn.Linear( + in_features=self.cross_attention_dim, + out_features=self.inner_dim, + bias_attr=bias, + ) + else: + self.to_k = None + self.to_v = None + if self.added_kv_proj_dim is not None: + self.add_k_proj = paddle.nn.Linear( + in_features=added_kv_proj_dim, out_features=self.inner_dim + ) + self.add_v_proj = paddle.nn.Linear( + in_features=added_kv_proj_dim, out_features=self.inner_dim + ) + if self.context_pre_only is not None: + self.add_q_proj = paddle.nn.Linear( + in_features=added_kv_proj_dim, out_features=self.inner_dim + ) + self.to_out = paddle.nn.LayerList(sublayers=[]) + self.to_out.append( + paddle.nn.Linear( + in_features=self.inner_dim, + out_features=self.out_dim, + bias_attr=out_bias, + ) + ) + self.to_out.append(paddle.nn.Dropout(p=dropout)) + if self.context_pre_only is not None and not self.context_pre_only: + self.to_add_out = paddle.nn.Linear( + in_features=self.inner_dim, + out_features=self.out_dim, + bias_attr=out_bias, + ) + if processor is None: + processor = ( + AttnProcessor2_0() +>>>>>> if hasattr(torch.nn.functional, "scaled_dot_product_attention") + and self.scale_qk + else AttnProcessor() + ) + self.set_processor(processor) + + def set_use_npu_flash_attention(self, use_npu_flash_attention: bool) -> None: + """ + Set whether to use npu flash attention from `torch_npu` or not. + + """ + if use_npu_flash_attention: + processor = AttnProcessorNPU() + else: + processor = ( + AttnProcessor2_0() +>>>>>> if hasattr(torch.nn.functional, "scaled_dot_product_attention") + and self.scale_qk + else AttnProcessor() + ) + self.set_processor(processor) + + def set_use_memory_efficient_attention_xformers( + self, + use_memory_efficient_attention_xformers: bool, + attention_op: Optional[Callable] = None, + ) -> None: + """ + Set whether to use memory efficient attention from `xformers` or not. + + Args: + use_memory_efficient_attention_xformers (`bool`): + Whether to use memory efficient attention from `xformers` or not. + attention_op (`Callable`, *optional*): + The attention operation to use. Defaults to `None` which uses the default attention operation from + `xformers`. + """ + is_lora = hasattr(self, "processor") and isinstance( + self.processor, LORA_ATTENTION_PROCESSORS + ) + is_custom_diffusion = hasattr(self, "processor") and isinstance( + self.processor, + ( + CustomDiffusionAttnProcessor, + CustomDiffusionXFormersAttnProcessor, + CustomDiffusionAttnProcessor2_0, + ), + ) + is_added_kv_processor = hasattr(self, "processor") and isinstance( + self.processor, + ( + AttnAddedKVProcessor, + AttnAddedKVProcessor2_0, + SlicedAttnAddedKVProcessor, + XFormersAttnAddedKVProcessor, + LoRAAttnAddedKVProcessor, + ), + ) + if use_memory_efficient_attention_xformers: + if is_added_kv_processor and (is_lora or is_custom_diffusion): + raise NotImplementedError( + f"Memory efficient attention is currently not supported for LoRA or custom diffusion for attention processor type {self.processor}" + ) + if not is_xformers_available(): + raise ModuleNotFoundError( + "Refer to https://github.com/facebookresearch/xformers for more information on how to install xformers", + name="xformers", + ) + elif not paddle.device.cuda.device_count() >= 1: + raise ValueError( + "torch.cuda.is_available() should be True but is False. xformers' memory efficient attention is only available for GPU " + ) + else: + try: + _ = xformers.ops.memory_efficient_attention( + paddle.randn((1, 2, 40), device="cuda"), + paddle.randn((1, 2, 40), device="cuda"), + paddle.randn((1, 2, 40), device="cuda"), + ) + except Exception as e: + raise e + if is_lora: + processor = LoRAXFormersAttnProcessor( + hidden_size=self.processor.hidden_size, + cross_attention_dim=self.processor.cross_attention_dim, + rank=self.processor.rank, + attention_op=attention_op, + ) + processor.set_state_dict(state_dict=self.processor.state_dict()) + processor.to(self.processor.to_q_lora.up.weight.place) + elif is_custom_diffusion: + processor = CustomDiffusionXFormersAttnProcessor( + train_kv=self.processor.train_kv, + train_q_out=self.processor.train_q_out, + hidden_size=self.processor.hidden_size, + cross_attention_dim=self.processor.cross_attention_dim, + attention_op=attention_op, + ) + processor.set_state_dict(state_dict=self.processor.state_dict()) + if hasattr(self.processor, "to_k_custom_diffusion"): + processor.to(self.processor.to_k_custom_diffusion.weight.place) + elif is_added_kv_processor: + logger.info( + "Memory efficient attention with `xformers` might currently not work correctly if an attention mask is required for the attention operation." + ) + processor = XFormersAttnAddedKVProcessor(attention_op=attention_op) + else: + processor = XFormersAttnProcessor(attention_op=attention_op) + elif is_lora: + attn_processor_class = ( + LoRAAttnProcessor2_0 +>>>>>> if hasattr(torch.nn.functional, "scaled_dot_product_attention") + else LoRAAttnProcessor + ) + processor = attn_processor_class( + hidden_size=self.processor.hidden_size, + cross_attention_dim=self.processor.cross_attention_dim, + rank=self.processor.rank, + ) + processor.set_state_dict(state_dict=self.processor.state_dict()) + processor.to(self.processor.to_q_lora.up.weight.place) + elif is_custom_diffusion: + attn_processor_class = ( + CustomDiffusionAttnProcessor2_0 +>>>>>> if hasattr(torch.nn.functional, "scaled_dot_product_attention") + else CustomDiffusionAttnProcessor + ) + processor = attn_processor_class( + train_kv=self.processor.train_kv, + train_q_out=self.processor.train_q_out, + hidden_size=self.processor.hidden_size, + cross_attention_dim=self.processor.cross_attention_dim, + ) + processor.set_state_dict(state_dict=self.processor.state_dict()) + if hasattr(self.processor, "to_k_custom_diffusion"): + processor.to(self.processor.to_k_custom_diffusion.weight.place) + else: + processor = ( + AttnProcessor2_0() +>>>>>> if hasattr(torch.nn.functional, "scaled_dot_product_attention") + and self.scale_qk + else AttnProcessor() + ) + self.set_processor(processor) + + def set_attention_slice(self, slice_size: int) -> None: + """ + Set the slice size for attention computation. + + Args: + slice_size (`int`): + The slice size for attention computation. + """ + if slice_size is not None and slice_size > self.sliceable_head_dim: + raise ValueError( + f"slice_size {slice_size} has to be smaller or equal to {self.sliceable_head_dim}." + ) + if slice_size is not None and self.added_kv_proj_dim is not None: + processor = SlicedAttnAddedKVProcessor(slice_size) + elif slice_size is not None: + processor = SlicedAttnProcessor(slice_size) + elif self.added_kv_proj_dim is not None: + processor = AttnAddedKVProcessor() + else: + processor = ( + AttnProcessor2_0() +>>>>>> if hasattr(torch.nn.functional, "scaled_dot_product_attention") + and self.scale_qk + else AttnProcessor() + ) + self.set_processor(processor) + + def set_processor(self, processor: "AttnProcessor") -> None: + """ + Set the attention processor to use. + + Args: + processor (`AttnProcessor`): + The attention processor to use. + """ + if ( + hasattr(self, "processor") + and isinstance(self.processor, paddle.nn.Layer) + and not isinstance(processor, paddle.nn.Layer) + ): + logger.info( + f"You are removing possibly trained weights of {self.processor} with {processor}" + ) + self._modules.pop("processor") + self.processor = processor + + def get_processor( + self, return_deprecated_lora: bool = False + ) -> "AttentionProcessor": + """ + Get the attention processor in use. + + Args: + return_deprecated_lora (`bool`, *optional*, defaults to `False`): + Set to `True` to return the deprecated LoRA attention processor. + + Returns: + "AttentionProcessor": The attention processor in use. + """ + if not return_deprecated_lora: + return self.processor + is_lora_activated = { + name: (module.lora_layer is not None) + for name, module in self.named_sublayers(include_self=True) + if hasattr(module, "lora_layer") + } + if not any(is_lora_activated.values()): + return self.processor + is_lora_activated.pop("add_k_proj", None) + is_lora_activated.pop("add_v_proj", None) + if not all(is_lora_activated.values()): + raise ValueError( + f"Make sure that either all layers or no layers have LoRA activated, but have {is_lora_activated}" + ) + non_lora_processor_cls_name = self.processor.__class__.__name__ + lora_processor_cls = getattr( + import_module(__name__), "LoRA" + non_lora_processor_cls_name + ) + hidden_size = self.inner_dim + if lora_processor_cls in [ + LoRAAttnProcessor, + LoRAAttnProcessor2_0, + LoRAXFormersAttnProcessor, + ]: + kwargs = { + "cross_attention_dim": self.cross_attention_dim, + "rank": self.to_q.lora_layer.rank, + "network_alpha": self.to_q.lora_layer.network_alpha, + "q_rank": self.to_q.lora_layer.rank, + "q_hidden_size": self.to_q.lora_layer.out_features, + "k_rank": self.to_k.lora_layer.rank, + "k_hidden_size": self.to_k.lora_layer.out_features, + "v_rank": self.to_v.lora_layer.rank, + "v_hidden_size": self.to_v.lora_layer.out_features, + "out_rank": self.to_out[0].lora_layer.rank, + "out_hidden_size": self.to_out[0].lora_layer.out_features, + } + if hasattr(self.processor, "attention_op"): + kwargs["attention_op"] = self.processor.attention_op + lora_processor = lora_processor_cls(hidden_size, **kwargs) + lora_processor.to_q_lora.load_state_dict(self.to_q.lora_layer.state_dict()) + lora_processor.to_k_lora.load_state_dict(self.to_k.lora_layer.state_dict()) + lora_processor.to_v_lora.load_state_dict(self.to_v.lora_layer.state_dict()) + lora_processor.to_out_lora.load_state_dict( + self.to_out[0].lora_layer.state_dict() + ) + elif lora_processor_cls == LoRAAttnAddedKVProcessor: + lora_processor = lora_processor_cls( + hidden_size, + cross_attention_dim=self.add_k_proj.weight.shape[0], + rank=self.to_q.lora_layer.rank, + network_alpha=self.to_q.lora_layer.network_alpha, + ) + lora_processor.to_q_lora.load_state_dict(self.to_q.lora_layer.state_dict()) + lora_processor.to_k_lora.load_state_dict(self.to_k.lora_layer.state_dict()) + lora_processor.to_v_lora.load_state_dict(self.to_v.lora_layer.state_dict()) + lora_processor.to_out_lora.load_state_dict( + self.to_out[0].lora_layer.state_dict() + ) + if self.add_k_proj.lora_layer is not None: + lora_processor.add_k_proj_lora.load_state_dict( + self.add_k_proj.lora_layer.state_dict() + ) + lora_processor.add_v_proj_lora.load_state_dict( + self.add_v_proj.lora_layer.state_dict() + ) + else: + lora_processor.add_k_proj_lora = None + lora_processor.add_v_proj_lora = None + else: + raise ValueError(f"{lora_processor_cls} does not exist.") + return lora_processor + + def forward( + self, + hidden_states: paddle.Tensor, + encoder_hidden_states: Optional[paddle.Tensor] = None, + attention_mask: Optional[paddle.Tensor] = None, + **cross_attention_kwargs, + ) -> paddle.Tensor: + """ + The forward method of the `Attention` class. + + Args: + hidden_states (`torch.Tensor`): + The hidden states of the query. + encoder_hidden_states (`torch.Tensor`, *optional*): + The hidden states of the encoder. + attention_mask (`torch.Tensor`, *optional*): + The attention mask to use. If `None`, no mask is applied. + **cross_attention_kwargs: + Additional keyword arguments to pass along to the cross attention. + + Returns: + `torch.Tensor`: The output of the attention layer. + """ + print("processor:", self.processor, "2" * 200) + attn_parameters = set( + inspect.signature(self.processor.__call__).parameters.keys() + ) + quiet_attn_parameters = {"ip_adapter_masks"} + unused_kwargs = [ + k + for k, _ in cross_attention_kwargs.items() + if k not in attn_parameters and k not in quiet_attn_parameters + ] + if len(unused_kwargs) > 0: + logger.warning( + f"cross_attention_kwargs {unused_kwargs} are not expected by {self.processor.__class__.__name__} and will be ignored." + ) + cross_attention_kwargs = { + k: w for k, w in cross_attention_kwargs.items() if k in attn_parameters + } + return self.processor( + self, + hidden_states, + encoder_hidden_states=encoder_hidden_states, + attention_mask=attention_mask, + **cross_attention_kwargs, + ) + + def batch_to_head_dim(self, tensor: paddle.Tensor) -> paddle.Tensor: + """ + Reshape the tensor from `[batch_size, seq_len, dim]` to `[batch_size // heads, seq_len, dim * heads]`. `heads` + is the number of heads initialized while constructing the `Attention` class. + + Args: + tensor (`torch.Tensor`): The tensor to reshape. + + Returns: + `torch.Tensor`: The reshaped tensor. + """ + head_size = self.heads + batch_size, seq_len, dim = tensor.shape + tensor = tensor.reshape(batch_size // head_size, head_size, seq_len, dim) + tensor = tensor.permute(0, 2, 1, 3).reshape( + batch_size // head_size, seq_len, dim * head_size + ) + return tensor + + def head_to_batch_dim( + self, tensor: paddle.Tensor, out_dim: int = 3 + ) -> paddle.Tensor: + """ + Reshape the tensor from `[batch_size, seq_len, dim]` to `[batch_size, seq_len, heads, dim // heads]` `heads` is + the number of heads initialized while constructing the `Attention` class. + + Args: + tensor (`torch.Tensor`): The tensor to reshape. + out_dim (`int`, *optional*, defaults to `3`): The output dimension of the tensor. If `3`, the tensor is + reshaped to `[batch_size * heads, seq_len, dim // heads]`. + + Returns: + `torch.Tensor`: The reshaped tensor. + """ + head_size = self.heads + if tensor.ndim == 3: + batch_size, seq_len, dim = tensor.shape + extra_dim = 1 + else: + batch_size, extra_dim, seq_len, dim = tensor.shape + tensor = tensor.reshape( + batch_size, seq_len * extra_dim, head_size, dim // head_size + ) + tensor = tensor.permute(0, 2, 1, 3) + if out_dim == 3: + tensor = tensor.reshape( + batch_size * head_size, seq_len * extra_dim, dim // head_size + ) + return tensor + + def get_attention_scores( + self, + query: paddle.Tensor, + key: paddle.Tensor, + attention_mask: paddle.Tensor = None, + ) -> paddle.Tensor: + """ + Compute the attention scores. + + Args: + query (`torch.Tensor`): The query tensor. + key (`torch.Tensor`): The key tensor. + attention_mask (`torch.Tensor`, *optional*): The attention mask to use. If `None`, no mask is applied. + + Returns: + `torch.Tensor`: The attention probabilities/scores. + """ + dtype = query.dtype + if self.upcast_attention: + query = query.float() + key = key.float() + if attention_mask is None: + baddbmm_input = paddle.empty( + query.shape[0], + query.shape[1], + key.shape[1], + dtype=query.dtype, + device=query.place, + ) + beta = 0 + else: + baddbmm_input = attention_mask + beta = 1 + attention_scores = paddle.baddbmm( + input=baddbmm_input, + x=query, + y=key.transpose(-1, -2), + beta=beta, + alpha=self.scale, + ) + del baddbmm_input + if self.upcast_softmax: + attention_scores = attention_scores.float() + attention_probs = attention_scores.softmax(dim=-1) + del attention_scores + attention_probs = attention_probs.to(dtype) + return attention_probs + + def prepare_attention_mask( + self, + attention_mask: paddle.Tensor, + target_length: int, + batch_size: int, + out_dim: int = 3, + ) -> paddle.Tensor: + """ + Prepare the attention mask for the attention computation. + + Args: + attention_mask (`torch.Tensor`): + The attention mask to prepare. + target_length (`int`): + The target length of the attention mask. This is the length of the attention mask after padding. + batch_size (`int`): + The batch size, which is used to repeat the attention mask. + out_dim (`int`, *optional*, defaults to `3`): + The output dimension of the attention mask. Can be either `3` or `4`. + + Returns: + `torch.Tensor`: The prepared attention mask. + """ + head_size = self.heads + if attention_mask is None: + return attention_mask + current_length: int = attention_mask.shape[-1] + if current_length != target_length: + if attention_mask.device.type == "mps": + padding_shape = ( + attention_mask.shape[0], + attention_mask.shape[1], + target_length, + ) + padding = paddle.zeros( + padding_shape, + dtype=attention_mask.dtype, + device=attention_mask.place, + ) + attention_mask = paddle.cat([attention_mask, padding], dim=2) + else: + attention_mask = paddle.compat.pad( + attention_mask, (0, target_length), value=0.0 + ) + if out_dim == 3: + if attention_mask.shape[0] < batch_size * head_size: + attention_mask = attention_mask.repeat_interleave(head_size, dim=0) + elif out_dim == 4: + attention_mask = attention_mask.unsqueeze(1) + attention_mask = attention_mask.repeat_interleave(head_size, dim=1) + return attention_mask + + def norm_encoder_hidden_states( + self, encoder_hidden_states: paddle.Tensor + ) -> paddle.Tensor: + """ + Normalize the encoder hidden states. Requires `self.norm_cross` to be specified when constructing the + `Attention` class. + + Args: + encoder_hidden_states (`torch.Tensor`): Hidden states of the encoder. + + Returns: + `torch.Tensor`: The normalized encoder hidden states. + """ + assert ( + self.norm_cross is not None + ), "self.norm_cross must be defined to call self.norm_encoder_hidden_states" + if isinstance(self.norm_cross, paddle.nn.LayerNorm): + encoder_hidden_states = self.norm_cross(encoder_hidden_states) + elif isinstance(self.norm_cross, paddle.nn.GroupNorm): + encoder_hidden_states = encoder_hidden_states.transpose(1, 2) + encoder_hidden_states = self.norm_cross(encoder_hidden_states) + encoder_hidden_states = encoder_hidden_states.transpose(1, 2) + else: + assert False + return encoder_hidden_states + + @paddle.no_grad() + def fuse_projections(self, fuse=True): + device = self.to_q.weight.data.place + dtype = self.to_q.weight.data.dtype + if not self.is_cross_attention: + concatenated_weights = paddle.cat( + [self.to_q.weight.data, self.to_k.weight.data, self.to_v.weight.data] + ) + in_features = concatenated_weights.shape[1] + out_features = concatenated_weights.shape[0] + self.to_qkv = paddle.nn.Linear( + in_features=in_features, + out_features=out_features, + bias_attr=self.use_bias, + ) + self.to_qkv.weight.copy_(concatenated_weights) + if self.use_bias: + concatenated_bias = paddle.cat( + [self.to_q.bias.data, self.to_k.bias.data, self.to_v.bias.data] + ) + self.to_qkv.bias.copy_(concatenated_bias) + else: + concatenated_weights = paddle.cat( + [self.to_k.weight.data, self.to_v.weight.data] + ) + in_features = concatenated_weights.shape[1] + out_features = concatenated_weights.shape[0] + self.to_kv = paddle.nn.Linear( + in_features=in_features, + out_features=out_features, + bias_attr=self.use_bias, + ) + self.to_kv.weight.copy_(concatenated_weights) + if self.use_bias: + concatenated_bias = paddle.cat( + [self.to_k.bias.data, self.to_v.bias.data] + ) + self.to_kv.bias.copy_(concatenated_bias) + self.fused_projections = fuse + + +class AttnProcessor: + """ + Default processor for performing attention-related computations. + """ + + def __call__( + self, + attn: Attention, + hidden_states: paddle.Tensor, + encoder_hidden_states: Optional[paddle.Tensor] = None, + attention_mask: Optional[paddle.Tensor] = None, + temb: Optional[paddle.Tensor] = None, + *args, + **kwargs, + ) -> paddle.Tensor: + if len(args) > 0 or kwargs.get("scale", None) is not None: + deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." + deprecate("scale", "1.0.0", deprecation_message) + residual = hidden_states + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + input_ndim = hidden_states.ndim + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view( + batch_size, channel, height * width + ).transpose(1, 2) + batch_size, sequence_length, _ = ( + hidden_states.shape + if encoder_hidden_states is None + else encoder_hidden_states.shape + ) + attention_mask = attn.prepare_attention_mask( + attention_mask, sequence_length, batch_size + ) + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose( + 1, 2 + ) + query = attn.to_q(hidden_states) + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states( + encoder_hidden_states + ) + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + query = attn.head_to_batch_dim(query) + key = attn.head_to_batch_dim(key) + value = attn.head_to_batch_dim(value) + attention_probs = attn.get_attention_scores(query, key, attention_mask) + hidden_states = paddle.bmm(attention_probs, value) + hidden_states = attn.batch_to_head_dim(hidden_states) + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape( + batch_size, channel, height, width + ) + if attn.residual_connection: + hidden_states = hidden_states + residual + hidden_states = hidden_states / attn.rescale_output_factor + return hidden_states + + +class CustomDiffusionAttnProcessor(paddle.nn.Layer): + """ + Processor for implementing attention for the Custom Diffusion method. + + Args: + train_kv (`bool`, defaults to `True`): + Whether to newly train the key and value matrices corresponding to the text features. + train_q_out (`bool`, defaults to `True`): + Whether to newly train query matrices corresponding to the latent image features. + hidden_size (`int`, *optional*, defaults to `None`): + The hidden size of the attention layer. + cross_attention_dim (`int`, *optional*, defaults to `None`): + The number of channels in the `encoder_hidden_states`. + out_bias (`bool`, defaults to `True`): + Whether to include the bias parameter in `train_q_out`. + dropout (`float`, *optional*, defaults to 0.0): + The dropout probability to use. + """ + + def __init__( + self, + train_kv: bool = True, + train_q_out: bool = True, + hidden_size: Optional[int] = None, + cross_attention_dim: Optional[int] = None, + out_bias: bool = True, + dropout: float = 0.0, + ): + super().__init__() + self.train_kv = train_kv + self.train_q_out = train_q_out + self.hidden_size = hidden_size + self.cross_attention_dim = cross_attention_dim + if self.train_kv: + self.to_k_custom_diffusion = paddle.nn.Linear( + in_features=cross_attention_dim or hidden_size, + out_features=hidden_size, + bias_attr=False, + ) + self.to_v_custom_diffusion = paddle.nn.Linear( + in_features=cross_attention_dim or hidden_size, + out_features=hidden_size, + bias_attr=False, + ) + if self.train_q_out: + self.to_q_custom_diffusion = paddle.nn.Linear( + in_features=hidden_size, out_features=hidden_size, bias_attr=False + ) + self.to_out_custom_diffusion = paddle.nn.LayerList(sublayers=[]) + self.to_out_custom_diffusion.append( + paddle.nn.Linear( + in_features=hidden_size, + out_features=hidden_size, + bias_attr=out_bias, + ) + ) + self.to_out_custom_diffusion.append(paddle.nn.Dropout(p=dropout)) + + def __call__( + self, + attn: Attention, + hidden_states: paddle.Tensor, + encoder_hidden_states: Optional[paddle.Tensor] = None, + attention_mask: Optional[paddle.Tensor] = None, + ) -> paddle.Tensor: + batch_size, sequence_length, _ = hidden_states.shape + attention_mask = attn.prepare_attention_mask( + attention_mask, sequence_length, batch_size + ) + if self.train_q_out: + query = self.to_q_custom_diffusion(hidden_states).to(attn.to_q.weight.dtype) + else: + query = attn.to_q(hidden_states.to(attn.to_q.weight.dtype)) + if encoder_hidden_states is None: + crossattn = False + encoder_hidden_states = hidden_states + else: + crossattn = True + if attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states( + encoder_hidden_states + ) + if self.train_kv: + key = self.to_k_custom_diffusion( + encoder_hidden_states.to(self.to_k_custom_diffusion.weight.dtype) + ) + value = self.to_v_custom_diffusion( + encoder_hidden_states.to(self.to_v_custom_diffusion.weight.dtype) + ) + key = key.to(attn.to_q.weight.dtype) + value = value.to(attn.to_q.weight.dtype) + else: + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + if crossattn: + detach = paddle.ones_like(key) + detach[:, :1, :] = detach[:, :1, :] * 0.0 + key = detach * key + (1 - detach) * key.detach() + value = detach * value + (1 - detach) * value.detach() + query = attn.head_to_batch_dim(query) + key = attn.head_to_batch_dim(key) + value = attn.head_to_batch_dim(value) + attention_probs = attn.get_attention_scores(query, key, attention_mask) + hidden_states = paddle.bmm(attention_probs, value) + hidden_states = attn.batch_to_head_dim(hidden_states) + if self.train_q_out: + hidden_states = self.to_out_custom_diffusion[0](hidden_states) + hidden_states = self.to_out_custom_diffusion[1](hidden_states) + else: + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + return hidden_states + + +class AttnAddedKVProcessor: + """ + Processor for performing attention-related computations with extra learnable key and value matrices for the text + encoder. + """ + + def __call__( + self, + attn: Attention, + hidden_states: paddle.Tensor, + encoder_hidden_states: Optional[paddle.Tensor] = None, + attention_mask: Optional[paddle.Tensor] = None, + *args, + **kwargs, + ) -> paddle.Tensor: + if len(args) > 0 or kwargs.get("scale", None) is not None: + deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." + deprecate("scale", "1.0.0", deprecation_message) + residual = hidden_states + hidden_states = hidden_states.view( + hidden_states.shape[0], hidden_states.shape[1], -1 + ).transpose(1, 2) + batch_size, sequence_length, _ = hidden_states.shape + attention_mask = attn.prepare_attention_mask( + attention_mask, sequence_length, batch_size + ) + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states( + encoder_hidden_states + ) + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + query = attn.to_q(hidden_states) + query = attn.head_to_batch_dim(query) + encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states) + encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states) + encoder_hidden_states_key_proj = attn.head_to_batch_dim( + encoder_hidden_states_key_proj + ) + encoder_hidden_states_value_proj = attn.head_to_batch_dim( + encoder_hidden_states_value_proj + ) + if not attn.only_cross_attention: + key = attn.to_k(hidden_states) + value = attn.to_v(hidden_states) + key = attn.head_to_batch_dim(key) + value = attn.head_to_batch_dim(value) + key = paddle.cat([encoder_hidden_states_key_proj, key], dim=1) + value = paddle.cat([encoder_hidden_states_value_proj, value], dim=1) + else: + key = encoder_hidden_states_key_proj + value = encoder_hidden_states_value_proj + attention_probs = attn.get_attention_scores(query, key, attention_mask) + hidden_states = paddle.bmm(attention_probs, value) + hidden_states = attn.batch_to_head_dim(hidden_states) + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + hidden_states = hidden_states.transpose(-1, -2).reshape(residual.shape) + hidden_states = hidden_states + residual + return hidden_states + + +class AttnAddedKVProcessor2_0: + """ + Processor for performing scaled dot-product attention (enabled by default if you're using PyTorch 2.0), with extra + learnable key and value matrices for the text encoder. + """ + + def __init__(self): +>>>>>> if not hasattr(torch.nn.functional, "scaled_dot_product_attention"): + raise ImportError( + "AttnAddedKVProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0." + ) + + def __call__( + self, + attn: Attention, + hidden_states: paddle.Tensor, + encoder_hidden_states: Optional[paddle.Tensor] = None, + attention_mask: Optional[paddle.Tensor] = None, + *args, + **kwargs, + ) -> paddle.Tensor: + if len(args) > 0 or kwargs.get("scale", None) is not None: + deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." + deprecate("scale", "1.0.0", deprecation_message) + residual = hidden_states + hidden_states = hidden_states.view( + hidden_states.shape[0], hidden_states.shape[1], -1 + ).transpose(1, 2) + batch_size, sequence_length, _ = hidden_states.shape + attention_mask = attn.prepare_attention_mask( + attention_mask, sequence_length, batch_size, out_dim=4 + ) + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states( + encoder_hidden_states + ) + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + query = attn.to_q(hidden_states) + query = attn.head_to_batch_dim(query, out_dim=4) + encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states) + encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states) + encoder_hidden_states_key_proj = attn.head_to_batch_dim( + encoder_hidden_states_key_proj, out_dim=4 + ) + encoder_hidden_states_value_proj = attn.head_to_batch_dim( + encoder_hidden_states_value_proj, out_dim=4 + ) + if not attn.only_cross_attention: + key = attn.to_k(hidden_states) + value = attn.to_v(hidden_states) + key = attn.head_to_batch_dim(key, out_dim=4) + value = attn.head_to_batch_dim(value, out_dim=4) + key = paddle.cat([encoder_hidden_states_key_proj, key], dim=2) + value = paddle.cat([encoder_hidden_states_value_proj, value], dim=2) + else: + key = encoder_hidden_states_key_proj + value = encoder_hidden_states_value_proj + hidden_states = paddle.nn.functional.scaled_dot_product_attention( + query.transpose([0, 2, 1, 3]), + key.transpose([0, 2, 1, 3]), + value.transpose([0, 2, 1, 3]), + attn_mask=attention_mask, + dropout_p=0.0, + is_causal=False, + ).transpose([0, 2, 1, 3]) + hidden_states = hidden_states.transpose(1, 2).reshape( + batch_size, -1, residual.shape[1] + ) + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + hidden_states = hidden_states.transpose(-1, -2).reshape(residual.shape) + hidden_states = hidden_states + residual + return hidden_states + + +class JointAttnProcessor2_0: + """Attention processor used typically in processing the SD3-like self-attention projections.""" + + def __init__(self): +>>>>>> if not hasattr(torch.nn.functional, "scaled_dot_product_attention"): + raise ImportError( + "AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0." + ) + + def __call__( + self, + attn: Attention, + hidden_states: paddle.FloatTensor, + encoder_hidden_states: paddle.FloatTensor = None, + attention_mask: Optional[paddle.FloatTensor] = None, + *args, + **kwargs, + ) -> paddle.FloatTensor: + residual = hidden_states + input_ndim = hidden_states.ndim + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view( + batch_size, channel, height * width + ).transpose(1, 2) + context_input_ndim = encoder_hidden_states.ndim + if context_input_ndim == 4: + batch_size, channel, height, width = encoder_hidden_states.shape + encoder_hidden_states = encoder_hidden_states.view( + batch_size, channel, height * width + ).transpose(1, 2) + batch_size = encoder_hidden_states.shape[0] + query = attn.to_q(hidden_states) + key = attn.to_k(hidden_states) + value = attn.to_v(hidden_states) + encoder_hidden_states_query_proj = attn.add_q_proj(encoder_hidden_states) + encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states) + encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states) + query = paddle.cat([query, encoder_hidden_states_query_proj], dim=1) + key = paddle.cat([key, encoder_hidden_states_key_proj], dim=1) + value = paddle.cat([value, encoder_hidden_states_value_proj], dim=1) + inner_dim = key.shape[-1] + head_dim = inner_dim // attn.heads + query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + hidden_states = ( + hidden_states + ) = paddle.nn.functional.scaled_dot_product_attention( + query.transpose([0, 2, 1, 3]), + key.transpose([0, 2, 1, 3]), + value.transpose([0, 2, 1, 3]), + dropout_p=0.0, + is_causal=False, + ).transpose( + [0, 2, 1, 3] + ) + hidden_states = hidden_states.transpose(1, 2).reshape( + batch_size, -1, attn.heads * head_dim + ) + hidden_states = hidden_states.to(query.dtype) + hidden_states, encoder_hidden_states = ( + hidden_states[:, : residual.shape[1]], + hidden_states[:, residual.shape[1] :], + ) + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + if not attn.context_pre_only: + encoder_hidden_states = attn.to_add_out(encoder_hidden_states) + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape( + batch_size, channel, height, width + ) + if context_input_ndim == 4: + encoder_hidden_states = encoder_hidden_states.transpose(-1, -2).reshape( + batch_size, channel, height, width + ) + return hidden_states, encoder_hidden_states + + +class FusedJointAttnProcessor2_0: + """Attention processor used typically in processing the SD3-like self-attention projections.""" + + def __init__(self): +>>>>>> if not hasattr(torch.nn.functional, "scaled_dot_product_attention"): + raise ImportError( + "AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0." + ) + + def __call__( + self, + attn: Attention, + hidden_states: paddle.FloatTensor, + encoder_hidden_states: paddle.FloatTensor = None, + attention_mask: Optional[paddle.FloatTensor] = None, + *args, + **kwargs, + ) -> paddle.FloatTensor: + residual = hidden_states + input_ndim = hidden_states.ndim + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view( + batch_size, channel, height * width + ).transpose(1, 2) + context_input_ndim = encoder_hidden_states.ndim + if context_input_ndim == 4: + batch_size, channel, height, width = encoder_hidden_states.shape + encoder_hidden_states = encoder_hidden_states.view( + batch_size, channel, height * width + ).transpose(1, 2) + batch_size = encoder_hidden_states.shape[0] + qkv = attn.to_qkv(hidden_states) + split_size = qkv.shape[-1] // 3 + query, key, value = paddle.compat.split(qkv, split_size, dim=-1) + encoder_qkv = attn.to_added_qkv(encoder_hidden_states) + split_size = encoder_qkv.shape[-1] // 3 + ( + encoder_hidden_states_query_proj, + encoder_hidden_states_key_proj, + encoder_hidden_states_value_proj, + ) = paddle.compat.split(encoder_qkv, split_size, dim=-1) + query = paddle.cat([query, encoder_hidden_states_query_proj], dim=1) + key = paddle.cat([key, encoder_hidden_states_key_proj], dim=1) + value = paddle.cat([value, encoder_hidden_states_value_proj], dim=1) + inner_dim = key.shape[-1] + head_dim = inner_dim // attn.heads + query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + hidden_states = ( + hidden_states + ) = paddle.nn.functional.scaled_dot_product_attention( + query.transpose([0, 2, 1, 3]), + key.transpose([0, 2, 1, 3]), + value.transpose([0, 2, 1, 3]), + dropout_p=0.0, + is_causal=False, + ).transpose( + [0, 2, 1, 3] + ) + hidden_states = hidden_states.transpose(1, 2).reshape( + batch_size, -1, attn.heads * head_dim + ) + hidden_states = hidden_states.to(query.dtype) + hidden_states, encoder_hidden_states = ( + hidden_states[:, : residual.shape[1]], + hidden_states[:, residual.shape[1] :], + ) + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + if not attn.context_pre_only: + encoder_hidden_states = attn.to_add_out(encoder_hidden_states) + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape( + batch_size, channel, height, width + ) + if context_input_ndim == 4: + encoder_hidden_states = encoder_hidden_states.transpose(-1, -2).reshape( + batch_size, channel, height, width + ) + return hidden_states, encoder_hidden_states + + +class XFormersAttnAddedKVProcessor: + """ + Processor for implementing memory efficient attention using xFormers. + + Args: + attention_op (`Callable`, *optional*, defaults to `None`): + The base + [operator](https://facebookresearch.github.io/xformers/components/ops.html#xformers.ops.AttentionOpBase) to + use as the attention operator. It is recommended to set to `None`, and allow xFormers to choose the best + operator. + """ + + def __init__(self, attention_op: Optional[Callable] = None): + self.attention_op = attention_op + + def __call__( + self, + attn: Attention, + hidden_states: paddle.Tensor, + encoder_hidden_states: Optional[paddle.Tensor] = None, + attention_mask: Optional[paddle.Tensor] = None, + ) -> paddle.Tensor: + residual = hidden_states + hidden_states = hidden_states.view( + hidden_states.shape[0], hidden_states.shape[1], -1 + ).transpose(1, 2) + batch_size, sequence_length, _ = hidden_states.shape + attention_mask = attn.prepare_attention_mask( + attention_mask, sequence_length, batch_size + ) + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states( + encoder_hidden_states + ) + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + query = attn.to_q(hidden_states) + query = attn.head_to_batch_dim(query) + encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states) + encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states) + encoder_hidden_states_key_proj = attn.head_to_batch_dim( + encoder_hidden_states_key_proj + ) + encoder_hidden_states_value_proj = attn.head_to_batch_dim( + encoder_hidden_states_value_proj + ) + if not attn.only_cross_attention: + key = attn.to_k(hidden_states) + value = attn.to_v(hidden_states) + key = attn.head_to_batch_dim(key) + value = attn.head_to_batch_dim(value) + key = paddle.cat([encoder_hidden_states_key_proj, key], dim=1) + value = paddle.cat([encoder_hidden_states_value_proj, value], dim=1) + else: + key = encoder_hidden_states_key_proj + value = encoder_hidden_states_value_proj + hidden_states = xformers.ops.memory_efficient_attention( + query, + key, + value, + attn_bias=attention_mask, + op=self.attention_op, + scale=attn.scale, + ) + hidden_states = hidden_states.to(query.dtype) + hidden_states = attn.batch_to_head_dim(hidden_states) + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + hidden_states = hidden_states.transpose(-1, -2).reshape(residual.shape) + hidden_states = hidden_states + residual + return hidden_states + + +class XFormersAttnProcessor: + """ + Processor for implementing memory efficient attention using xFormers. + + Args: + attention_op (`Callable`, *optional*, defaults to `None`): + The base + [operator](https://facebookresearch.github.io/xformers/components/ops.html#xformers.ops.AttentionOpBase) to + use as the attention operator. It is recommended to set to `None`, and allow xFormers to choose the best + operator. + """ + + def __init__(self, attention_op: Optional[Callable] = None): + self.attention_op = attention_op + + def __call__( + self, + attn: Attention, + hidden_states: paddle.Tensor, + encoder_hidden_states: Optional[paddle.Tensor] = None, + attention_mask: Optional[paddle.Tensor] = None, + temb: Optional[paddle.Tensor] = None, + *args, + **kwargs, + ) -> paddle.Tensor: + if len(args) > 0 or kwargs.get("scale", None) is not None: + deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." + deprecate("scale", "1.0.0", deprecation_message) + residual = hidden_states + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + input_ndim = hidden_states.ndim + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view( + batch_size, channel, height * width + ).transpose(1, 2) + batch_size, key_tokens, _ = ( + hidden_states.shape + if encoder_hidden_states is None + else encoder_hidden_states.shape + ) + attention_mask = attn.prepare_attention_mask( + attention_mask, key_tokens, batch_size + ) + if attention_mask is not None: + _, query_tokens, _ = hidden_states.shape + attention_mask = attention_mask.expand(-1, query_tokens, -1) + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose( + 1, 2 + ) + query = attn.to_q(hidden_states) + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states( + encoder_hidden_states + ) + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + query = attn.head_to_batch_dim(query).contiguous() + key = attn.head_to_batch_dim(key).contiguous() + value = attn.head_to_batch_dim(value).contiguous() + hidden_states = xformers.ops.memory_efficient_attention( + query, + key, + value, + attn_bias=attention_mask, + op=self.attention_op, + scale=attn.scale, + ) + hidden_states = hidden_states.to(query.dtype) + hidden_states = attn.batch_to_head_dim(hidden_states) + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape( + batch_size, channel, height, width + ) + if attn.residual_connection: + hidden_states = hidden_states + residual + hidden_states = hidden_states / attn.rescale_output_factor + return hidden_states + + +class AttnProcessorNPU: + """ + Processor for implementing flash attention using torch_npu. Torch_npu supports only fp16 and bf16 data types. If + fp32 is used, F.scaled_dot_product_attention will be used for computation, but the acceleration effect on NPU is + not significant. + + """ + + def __init__(self): + if not is_torch_npu_available(): + raise ImportError( + "AttnProcessorNPU requires torch_npu extensions and is supported only on npu devices." + ) + + def __call__( + self, + attn: Attention, + hidden_states: paddle.Tensor, + encoder_hidden_states: Optional[paddle.Tensor] = None, + attention_mask: Optional[paddle.Tensor] = None, + temb: Optional[paddle.Tensor] = None, + *args, + **kwargs, + ) -> paddle.Tensor: + if len(args) > 0 or kwargs.get("scale", None) is not None: + deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." + deprecate("scale", "1.0.0", deprecation_message) + residual = hidden_states + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + input_ndim = hidden_states.ndim + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view( + batch_size, channel, height * width + ).transpose(1, 2) + batch_size, sequence_length, _ = ( + hidden_states.shape + if encoder_hidden_states is None + else encoder_hidden_states.shape + ) + if attention_mask is not None: + attention_mask = attn.prepare_attention_mask( + attention_mask, sequence_length, batch_size + ) + attention_mask = attention_mask.view( + batch_size, attn.heads, -1, attention_mask.shape[-1] + ) + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose( + 1, 2 + ) + query = attn.to_q(hidden_states) + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states( + encoder_hidden_states + ) + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + inner_dim = key.shape[-1] + head_dim = inner_dim // attn.heads + query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + if query.dtype in (paddle.float16, paddle.bfloat16): + hidden_states = torch_npu.npu_fusion_attention( + query, + key, + value, + attn.heads, + input_layout="BNSD", + pse=None, + atten_mask=attention_mask, + scale=1.0 / math.sqrt(query.shape[-1]), + pre_tockens=65536, + next_tockens=65536, + keep_prob=1.0, + sync=False, + inner_precise=0, + )[0] + else: + hidden_states = paddle.nn.functional.scaled_dot_product_attention( + query.transpose([0, 2, 1, 3]), + key.transpose([0, 2, 1, 3]), + value.transpose([0, 2, 1, 3]), + attn_mask=attention_mask, + dropout_p=0.0, + is_causal=False, + ).transpose([0, 2, 1, 3]) + hidden_states = hidden_states.transpose(1, 2).reshape( + batch_size, -1, attn.heads * head_dim + ) + hidden_states = hidden_states.to(query.dtype) + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape( + batch_size, channel, height, width + ) + if attn.residual_connection: + hidden_states = hidden_states + residual + hidden_states = hidden_states / attn.rescale_output_factor + return hidden_states + + +class AttnProcessor2_0: + """ + Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). + """ + + def __init__(self): + return + + def __call__( + self, + attn: Attention, + hidden_states: paddle.Tensor, + encoder_hidden_states: Optional[paddle.Tensor] = None, + attention_mask: Optional[paddle.Tensor] = None, + temb: Optional[paddle.Tensor] = None, + *args, + **kwargs, + ) -> paddle.Tensor: + residual = hidden_states + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + input_ndim = hidden_states.ndim + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view( + batch_size, channel, height * width + ).transpose(1, 2) + batch_size, sequence_length, _ = ( + hidden_states.shape + if encoder_hidden_states is None + else encoder_hidden_states.shape + ) + if attention_mask is not None: + attention_mask = attn.prepare_attention_mask( + attention_mask, sequence_length, batch_size + ) + attention_mask = attention_mask.view( + batch_size, attn.heads, -1, attention_mask.shape[-1] + ) + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose( + 1, 2 + ) + query = attn.to_q(hidden_states) + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states( + encoder_hidden_states + ) + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + inner_dim = key.shape[-1] + head_dim = inner_dim // attn.heads + query = paddle.transpose(query.view(batch_size, -1, attn.heads, head_dim),perm = [0,2,1]) + key = paddle.transpose(key.view(batch_size, -1, attn.heads, head_dim),perm = [0,2,1]) + value =paddle.transpose( value.view(batch_size, -1, attn.heads, head_dim),perm = [0,2,1]) + hidden_states = paddle.nn.functional.scaled_dot_product_attention( + query.transpose([0, 2, 1, 3]), + key.transpose([0, 2, 1, 3]), + value.transpose([0, 2, 1, 3]), + attn_mask=attention_mask, + dropout_p=0.0, + is_causal=False, + ).transpose([0, 2, 1, 3]) + hidden_states = hidden_states.transpose(1, 2).reshape( + batch_size, -1, attn.heads * head_dim + ) + hidden_states = hidden_states.to(query.dtype) + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape( + batch_size, channel, height, width + ) + if attn.residual_connection: + hidden_states = hidden_states + residual + hidden_states = hidden_states / attn.rescale_output_factor + return hidden_states + + +class HunyuanAttnProcessor2_0: + """ + Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). This is + used in the HunyuanDiT model. It applies a s normalization layer and rotary embedding on query and key vector. + """ + + def __init__(self): +>>>>>> if not hasattr(torch.nn.functional, "scaled_dot_product_attention"): + raise ImportError( + "AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0." + ) + + def __call__( + self, + attn: Attention, + hidden_states: paddle.Tensor, + encoder_hidden_states: Optional[paddle.Tensor] = None, + attention_mask: Optional[paddle.Tensor] = None, + temb: Optional[paddle.Tensor] = None, + image_rotary_emb: Optional[paddle.Tensor] = None, + ) -> paddle.Tensor: + from .embeddings import apply_rotary_emb + + residual = hidden_states + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + input_ndim = hidden_states.ndim + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view( + batch_size, channel, height * width + ).transpose(1, 2) + batch_size, sequence_length, _ = ( + hidden_states.shape + if encoder_hidden_states is None + else encoder_hidden_states.shape + ) + if attention_mask is not None: + attention_mask = attn.prepare_attention_mask( + attention_mask, sequence_length, batch_size + ) + attention_mask = attention_mask.view( + batch_size, attn.heads, -1, attention_mask.shape[-1] + ) + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose( + 1, 2 + ) + query = attn.to_q(hidden_states) + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states( + encoder_hidden_states + ) + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + inner_dim = key.shape[-1] + head_dim = inner_dim // attn.heads + query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + if attn.norm_q is not None: + query = attn.norm_q(query) + if attn.norm_k is not None: + key = attn.norm_k(key) + if image_rotary_emb is not None: + query = apply_rotary_emb(query, image_rotary_emb) + if not attn.is_cross_attention: + key = apply_rotary_emb(key, image_rotary_emb) + hidden_states = paddle.nn.functional.scaled_dot_product_attention( + query.transpose([0, 2, 1, 3]), + key.transpose([0, 2, 1, 3]), + value.transpose([0, 2, 1, 3]), + attn_mask=attention_mask, + dropout_p=0.0, + is_causal=False, + ).transpose([0, 2, 1, 3]) + hidden_states = hidden_states.transpose(1, 2).reshape( + batch_size, -1, attn.heads * head_dim + ) + hidden_states = hidden_states.to(query.dtype) + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape( + batch_size, channel, height, width + ) + if attn.residual_connection: + hidden_states = hidden_states + residual + hidden_states = hidden_states / attn.rescale_output_factor + return hidden_states + + +class FusedAttnProcessor2_0: + """ + Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). It uses + fused projection layers. For self-attention modules, all projection matrices (i.e., query, key, value) are fused. + For cross-attention modules, key and value projection matrices are fused. + + + + This API is currently 🧪 experimental in nature and can change in future. + + + """ + + def __init__(self): +>>>>>> if not hasattr(torch.nn.functional, "scaled_dot_product_attention"): + raise ImportError( + "FusedAttnProcessor2_0 requires at least PyTorch 2.0, to use it. Please upgrade PyTorch to > 2.0." + ) + + def __call__( + self, + attn: Attention, + hidden_states: paddle.Tensor, + encoder_hidden_states: Optional[paddle.Tensor] = None, + attention_mask: Optional[paddle.Tensor] = None, + temb: Optional[paddle.Tensor] = None, + *args, + **kwargs, + ) -> paddle.Tensor: + if len(args) > 0 or kwargs.get("scale", None) is not None: + deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." + deprecate("scale", "1.0.0", deprecation_message) + residual = hidden_states + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + input_ndim = hidden_states.ndim + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view( + batch_size, channel, height * width + ).transpose(1, 2) + batch_size, sequence_length, _ = ( + hidden_states.shape + if encoder_hidden_states is None + else encoder_hidden_states.shape + ) + if attention_mask is not None: + attention_mask = attn.prepare_attention_mask( + attention_mask, sequence_length, batch_size + ) + attention_mask = attention_mask.view( + batch_size, attn.heads, -1, attention_mask.shape[-1] + ) + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose( + 1, 2 + ) + if encoder_hidden_states is None: + qkv = attn.to_qkv(hidden_states) + split_size = qkv.shape[-1] // 3 + query, key, value = paddle.compat.split(qkv, split_size, dim=-1) + else: + if attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states( + encoder_hidden_states + ) + query = attn.to_q(hidden_states) + kv = attn.to_kv(encoder_hidden_states) + split_size = kv.shape[-1] // 2 + key, value = paddle.compat.split(kv, split_size, dim=-1) + inner_dim = key.shape[-1] + head_dim = inner_dim // attn.heads + query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + hidden_states = paddle.nn.functional.scaled_dot_product_attention( + query.transpose([0, 2, 1, 3]), + key.transpose([0, 2, 1, 3]), + value.transpose([0, 2, 1, 3]), + attn_mask=attention_mask, + dropout_p=0.0, + is_causal=False, + ).transpose([0, 2, 1, 3]) + hidden_states = hidden_states.transpose(1, 2).reshape( + batch_size, -1, attn.heads * head_dim + ) + hidden_states = hidden_states.to(query.dtype) + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape( + batch_size, channel, height, width + ) + if attn.residual_connection: + hidden_states = hidden_states + residual + hidden_states = hidden_states / attn.rescale_output_factor + return hidden_states + + +class CustomDiffusionXFormersAttnProcessor(paddle.nn.Layer): + """ + Processor for implementing memory efficient attention using xFormers for the Custom Diffusion method. + + Args: + train_kv (`bool`, defaults to `True`): + Whether to newly train the key and value matrices corresponding to the text features. + train_q_out (`bool`, defaults to `True`): + Whether to newly train query matrices corresponding to the latent image features. + hidden_size (`int`, *optional*, defaults to `None`): + The hidden size of the attention layer. + cross_attention_dim (`int`, *optional*, defaults to `None`): + The number of channels in the `encoder_hidden_states`. + out_bias (`bool`, defaults to `True`): + Whether to include the bias parameter in `train_q_out`. + dropout (`float`, *optional*, defaults to 0.0): + The dropout probability to use. + attention_op (`Callable`, *optional*, defaults to `None`): + The base + [operator](https://facebookresearch.github.io/xformers/components/ops.html#xformers.ops.AttentionOpBase) to use + as the attention operator. It is recommended to set to `None`, and allow xFormers to choose the best operator. + """ + + def __init__( + self, + train_kv: bool = True, + train_q_out: bool = False, + hidden_size: Optional[int] = None, + cross_attention_dim: Optional[int] = None, + out_bias: bool = True, + dropout: float = 0.0, + attention_op: Optional[Callable] = None, + ): + super().__init__() + self.train_kv = train_kv + self.train_q_out = train_q_out + self.hidden_size = hidden_size + self.cross_attention_dim = cross_attention_dim + self.attention_op = attention_op + if self.train_kv: + self.to_k_custom_diffusion = paddle.nn.Linear( + in_features=cross_attention_dim or hidden_size, + out_features=hidden_size, + bias_attr=False, + ) + self.to_v_custom_diffusion = paddle.nn.Linear( + in_features=cross_attention_dim or hidden_size, + out_features=hidden_size, + bias_attr=False, + ) + if self.train_q_out: + self.to_q_custom_diffusion = paddle.nn.Linear( + in_features=hidden_size, out_features=hidden_size, bias_attr=False + ) + self.to_out_custom_diffusion = paddle.nn.LayerList(sublayers=[]) + self.to_out_custom_diffusion.append( + paddle.nn.Linear( + in_features=hidden_size, + out_features=hidden_size, + bias_attr=out_bias, + ) + ) + self.to_out_custom_diffusion.append(paddle.nn.Dropout(p=dropout)) + + def __call__( + self, + attn: Attention, + hidden_states: paddle.Tensor, + encoder_hidden_states: Optional[paddle.Tensor] = None, + attention_mask: Optional[paddle.Tensor] = None, + ) -> paddle.Tensor: + batch_size, sequence_length, _ = ( + hidden_states.shape + if encoder_hidden_states is None + else encoder_hidden_states.shape + ) + attention_mask = attn.prepare_attention_mask( + attention_mask, sequence_length, batch_size + ) + if self.train_q_out: + query = self.to_q_custom_diffusion(hidden_states).to(attn.to_q.weight.dtype) + else: + query = attn.to_q(hidden_states.to(attn.to_q.weight.dtype)) + if encoder_hidden_states is None: + crossattn = False + encoder_hidden_states = hidden_states + else: + crossattn = True + if attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states( + encoder_hidden_states + ) + if self.train_kv: + key = self.to_k_custom_diffusion( + encoder_hidden_states.to(self.to_k_custom_diffusion.weight.dtype) + ) + value = self.to_v_custom_diffusion( + encoder_hidden_states.to(self.to_v_custom_diffusion.weight.dtype) + ) + key = key.to(attn.to_q.weight.dtype) + value = value.to(attn.to_q.weight.dtype) + else: + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + if crossattn: + detach = paddle.ones_like(key) + detach[:, :1, :] = detach[:, :1, :] * 0.0 + key = detach * key + (1 - detach) * key.detach() + value = detach * value + (1 - detach) * value.detach() + query = attn.head_to_batch_dim(query).contiguous() + key = attn.head_to_batch_dim(key).contiguous() + value = attn.head_to_batch_dim(value).contiguous() + hidden_states = xformers.ops.memory_efficient_attention( + query, + key, + value, + attn_bias=attention_mask, + op=self.attention_op, + scale=attn.scale, + ) + hidden_states = hidden_states.to(query.dtype) + hidden_states = attn.batch_to_head_dim(hidden_states) + if self.train_q_out: + hidden_states = self.to_out_custom_diffusion[0](hidden_states) + hidden_states = self.to_out_custom_diffusion[1](hidden_states) + else: + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + return hidden_states + + +class CustomDiffusionAttnProcessor2_0(paddle.nn.Layer): + """ + Processor for implementing attention for the Custom Diffusion method using PyTorch 2.0’s memory-efficient scaled + dot-product attention. + + Args: + train_kv (`bool`, defaults to `True`): + Whether to newly train the key and value matrices corresponding to the text features. + train_q_out (`bool`, defaults to `True`): + Whether to newly train query matrices corresponding to the latent image features. + hidden_size (`int`, *optional*, defaults to `None`): + The hidden size of the attention layer. + cross_attention_dim (`int`, *optional*, defaults to `None`): + The number of channels in the `encoder_hidden_states`. + out_bias (`bool`, defaults to `True`): + Whether to include the bias parameter in `train_q_out`. + dropout (`float`, *optional*, defaults to 0.0): + The dropout probability to use. + """ + + def __init__( + self, + train_kv: bool = True, + train_q_out: bool = True, + hidden_size: Optional[int] = None, + cross_attention_dim: Optional[int] = None, + out_bias: bool = True, + dropout: float = 0.0, + ): + super().__init__() + self.train_kv = train_kv + self.train_q_out = train_q_out + self.hidden_size = hidden_size + self.cross_attention_dim = cross_attention_dim + if self.train_kv: + self.to_k_custom_diffusion = paddle.nn.Linear( + in_features=cross_attention_dim or hidden_size, + out_features=hidden_size, + bias_attr=False, + ) + self.to_v_custom_diffusion = paddle.nn.Linear( + in_features=cross_attention_dim or hidden_size, + out_features=hidden_size, + bias_attr=False, + ) + if self.train_q_out: + self.to_q_custom_diffusion = paddle.nn.Linear( + in_features=hidden_size, out_features=hidden_size, bias_attr=False + ) + self.to_out_custom_diffusion = paddle.nn.LayerList(sublayers=[]) + self.to_out_custom_diffusion.append( + paddle.nn.Linear( + in_features=hidden_size, + out_features=hidden_size, + bias_attr=out_bias, + ) + ) + self.to_out_custom_diffusion.append(paddle.nn.Dropout(p=dropout)) + + def __call__( + self, + attn: Attention, + hidden_states: paddle.Tensor, + encoder_hidden_states: Optional[paddle.Tensor] = None, + attention_mask: Optional[paddle.Tensor] = None, + ) -> paddle.Tensor: + batch_size, sequence_length, _ = hidden_states.shape + attention_mask = attn.prepare_attention_mask( + attention_mask, sequence_length, batch_size + ) + if self.train_q_out: + query = self.to_q_custom_diffusion(hidden_states) + else: + query = attn.to_q(hidden_states) + if encoder_hidden_states is None: + crossattn = False + encoder_hidden_states = hidden_states + else: + crossattn = True + if attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states( + encoder_hidden_states + ) + if self.train_kv: + key = self.to_k_custom_diffusion( + encoder_hidden_states.to(self.to_k_custom_diffusion.weight.dtype) + ) + value = self.to_v_custom_diffusion( + encoder_hidden_states.to(self.to_v_custom_diffusion.weight.dtype) + ) + key = key.to(attn.to_q.weight.dtype) + value = value.to(attn.to_q.weight.dtype) + else: + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + if crossattn: + detach = paddle.ones_like(key) + detach[:, :1, :] = detach[:, :1, :] * 0.0 + key = detach * key + (1 - detach) * key.detach() + value = detach * value + (1 - detach) * value.detach() + inner_dim = hidden_states.shape[-1] + head_dim = inner_dim // attn.heads + query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + hidden_states = paddle.nn.functional.scaled_dot_product_attention( + query.transpose([0, 2, 1, 3]), + key.transpose([0, 2, 1, 3]), + value.transpose([0, 2, 1, 3]), + attn_mask=attention_mask, + dropout_p=0.0, + is_causal=False, + ).transpose([0, 2, 1, 3]) + hidden_states = hidden_states.transpose(1, 2).reshape( + batch_size, -1, attn.heads * head_dim + ) + hidden_states = hidden_states.to(query.dtype) + if self.train_q_out: + hidden_states = self.to_out_custom_diffusion[0](hidden_states) + hidden_states = self.to_out_custom_diffusion[1](hidden_states) + else: + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + return hidden_states + + +class SlicedAttnProcessor: + """ + Processor for implementing sliced attention. + + Args: + slice_size (`int`, *optional*): + The number of steps to compute attention. Uses as many slices as `attention_head_dim // slice_size`, and + `attention_head_dim` must be a multiple of the `slice_size`. + """ + + def __init__(self, slice_size: int): + self.slice_size = slice_size + + def __call__( + self, + attn: Attention, + hidden_states: paddle.Tensor, + encoder_hidden_states: Optional[paddle.Tensor] = None, + attention_mask: Optional[paddle.Tensor] = None, + ) -> paddle.Tensor: + residual = hidden_states + input_ndim = hidden_states.ndim + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view( + batch_size, channel, height * width + ).transpose(1, 2) + batch_size, sequence_length, _ = ( + hidden_states.shape + if encoder_hidden_states is None + else encoder_hidden_states.shape + ) + attention_mask = attn.prepare_attention_mask( + attention_mask, sequence_length, batch_size + ) + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose( + 1, 2 + ) + query = attn.to_q(hidden_states) + dim = query.shape[-1] + query = attn.head_to_batch_dim(query) + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states( + encoder_hidden_states + ) + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + key = attn.head_to_batch_dim(key) + value = attn.head_to_batch_dim(value) + batch_size_attention, query_tokens, _ = query.shape + hidden_states = paddle.zeros( + (batch_size_attention, query_tokens, dim // attn.heads), + device=query.place, + dtype=query.dtype, + ) + for i in range(batch_size_attention // self.slice_size): + start_idx = i * self.slice_size + end_idx = (i + 1) * self.slice_size + query_slice = query[start_idx:end_idx] + key_slice = key[start_idx:end_idx] + attn_mask_slice = ( + attention_mask[start_idx:end_idx] + if attention_mask is not None + else None + ) + attn_slice = attn.get_attention_scores( + query_slice, key_slice, attn_mask_slice + ) + attn_slice = paddle.bmm(attn_slice, value[start_idx:end_idx]) + hidden_states[start_idx:end_idx] = attn_slice + hidden_states = attn.batch_to_head_dim(hidden_states) + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape( + batch_size, channel, height, width + ) + if attn.residual_connection: + hidden_states = hidden_states + residual + hidden_states = hidden_states / attn.rescale_output_factor + return hidden_states + + +class SlicedAttnAddedKVProcessor: + """ + Processor for implementing sliced attention with extra learnable key and value matrices for the text encoder. + + Args: + slice_size (`int`, *optional*): + The number of steps to compute attention. Uses as many slices as `attention_head_dim // slice_size`, and + `attention_head_dim` must be a multiple of the `slice_size`. + """ + + def __init__(self, slice_size): + self.slice_size = slice_size + + def __call__( + self, + attn: "Attention", + hidden_states: paddle.Tensor, + encoder_hidden_states: Optional[paddle.Tensor] = None, + attention_mask: Optional[paddle.Tensor] = None, + temb: Optional[paddle.Tensor] = None, + ) -> paddle.Tensor: + residual = hidden_states + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + hidden_states = hidden_states.view( + hidden_states.shape[0], hidden_states.shape[1], -1 + ).transpose(1, 2) + batch_size, sequence_length, _ = hidden_states.shape + attention_mask = attn.prepare_attention_mask( + attention_mask, sequence_length, batch_size + ) + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states( + encoder_hidden_states + ) + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + query = attn.to_q(hidden_states) + dim = query.shape[-1] + query = attn.head_to_batch_dim(query) + encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states) + encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states) + encoder_hidden_states_key_proj = attn.head_to_batch_dim( + encoder_hidden_states_key_proj + ) + encoder_hidden_states_value_proj = attn.head_to_batch_dim( + encoder_hidden_states_value_proj + ) + if not attn.only_cross_attention: + key = attn.to_k(hidden_states) + value = attn.to_v(hidden_states) + key = attn.head_to_batch_dim(key) + value = attn.head_to_batch_dim(value) + key = paddle.cat([encoder_hidden_states_key_proj, key], dim=1) + value = paddle.cat([encoder_hidden_states_value_proj, value], dim=1) + else: + key = encoder_hidden_states_key_proj + value = encoder_hidden_states_value_proj + batch_size_attention, query_tokens, _ = query.shape + hidden_states = paddle.zeros( + (batch_size_attention, query_tokens, dim // attn.heads), + device=query.place, + dtype=query.dtype, + ) + for i in range(batch_size_attention // self.slice_size): + start_idx = i * self.slice_size + end_idx = (i + 1) * self.slice_size + query_slice = query[start_idx:end_idx] + key_slice = key[start_idx:end_idx] + attn_mask_slice = ( + attention_mask[start_idx:end_idx] + if attention_mask is not None + else None + ) + attn_slice = attn.get_attention_scores( + query_slice, key_slice, attn_mask_slice + ) + attn_slice = paddle.bmm(attn_slice, value[start_idx:end_idx]) + hidden_states[start_idx:end_idx] = attn_slice + hidden_states = attn.batch_to_head_dim(hidden_states) + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + hidden_states = hidden_states.transpose(-1, -2).reshape(residual.shape) + hidden_states = hidden_states + residual + return hidden_states + + +class SpatialNorm(paddle.nn.Layer): + """ + Spatially conditioned normalization as defined in https://arxiv.org/abs/2209.09002. + + Args: + f_channels (`int`): + The number of channels for input to group normalization layer, and output of the spatial norm layer. + zq_channels (`int`): + The number of channels for the quantized vector as described in the paper. + """ + + def __init__(self, f_channels: int, zq_channels: int): + super().__init__() + self.norm_layer = paddle.nn.GroupNorm( + num_channels=f_channels, + num_groups=32, + epsilon=1e-06, + weight_attr=True, + bias_attr=True, + ) + self.conv_y = paddle.nn.Conv2d( + zq_channels, f_channels, kernel_size=1, stride=1, padding=0 + ) + self.conv_b = paddle.nn.Conv2d( + zq_channels, f_channels, kernel_size=1, stride=1, padding=0 + ) + + def forward(self, f: paddle.Tensor, zq: paddle.Tensor) -> paddle.Tensor: + f_size = f.shape[-2:] + zq = paddle.nn.functional.interpolate(x=zq, size=f_size, mode="nearest") + norm_f = self.norm_layer(f) + new_f = norm_f * self.conv_y(zq) + self.conv_b(zq) + return new_f + + +class LoRAAttnProcessor(paddle.nn.Layer): + def __init__( + self, + hidden_size: int, + cross_attention_dim: Optional[int] = None, + rank: int = 4, + network_alpha: Optional[int] = None, + **kwargs, + ): + deprecation_message = "Using LoRAAttnProcessor is deprecated. Please use the PEFT backend for all things LoRA. You can install PEFT by running `pip install peft`." + deprecate( + "LoRAAttnProcessor", "0.30.0", deprecation_message, standard_warn=False + ) + super().__init__() + self.hidden_size = hidden_size + self.cross_attention_dim = cross_attention_dim + self.rank = rank + q_rank = kwargs.pop("q_rank", None) + q_hidden_size = kwargs.pop("q_hidden_size", None) + q_rank = q_rank if q_rank is not None else rank + q_hidden_size = q_hidden_size if q_hidden_size is not None else hidden_size + v_rank = kwargs.pop("v_rank", None) + v_hidden_size = kwargs.pop("v_hidden_size", None) + v_rank = v_rank if v_rank is not None else rank + v_hidden_size = v_hidden_size if v_hidden_size is not None else hidden_size + out_rank = kwargs.pop("out_rank", None) + out_hidden_size = kwargs.pop("out_hidden_size", None) + out_rank = out_rank if out_rank is not None else rank + out_hidden_size = ( + out_hidden_size if out_hidden_size is not None else hidden_size + ) + self.to_q_lora = LoRALinearLayer( + q_hidden_size, q_hidden_size, q_rank, network_alpha + ) + self.to_k_lora = LoRALinearLayer( + cross_attention_dim or hidden_size, hidden_size, rank, network_alpha + ) + self.to_v_lora = LoRALinearLayer( + cross_attention_dim or v_hidden_size, v_hidden_size, v_rank, network_alpha + ) + self.to_out_lora = LoRALinearLayer( + out_hidden_size, out_hidden_size, out_rank, network_alpha + ) + + def __call__( + self, attn: Attention, hidden_states: paddle.Tensor, **kwargs + ) -> paddle.Tensor: + self_cls_name = self.__class__.__name__ + deprecate( + self_cls_name, + "0.26.0", + f"Make sure use {self_cls_name[4:]} instead by settingLoRA layers to `self.{{to_q,to_k,to_v,to_out[0]}}.lora_layer` respectively. This will be done automatically when using `LoraLoaderMixin.load_lora_weights`", + ) + attn.to_q.lora_layer = self.to_q_lora.to(hidden_states.place) + attn.to_k.lora_layer = self.to_k_lora.to(hidden_states.place) + attn.to_v.lora_layer = self.to_v_lora.to(hidden_states.place) + attn.to_out[0].lora_layer = self.to_out_lora.to(hidden_states.place) + attn._modules.pop("processor") + attn.processor = AttnProcessor() + return attn.processor(attn, hidden_states, **kwargs) + + +class LoRAAttnProcessor2_0(paddle.nn.Layer): + def __init__( + self, + hidden_size: int, + cross_attention_dim: Optional[int] = None, + rank: int = 4, + network_alpha: Optional[int] = None, + **kwargs, + ): + deprecation_message = "Using LoRAAttnProcessor is deprecated. Please use the PEFT backend for all things LoRA. You can install PEFT by running `pip install peft`." + deprecate( + "LoRAAttnProcessor2_0", "0.30.0", deprecation_message, standard_warn=False + ) + super().__init__() +>>>>>> if not hasattr(torch.nn.functional, "scaled_dot_product_attention"): + raise ImportError( + "AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0." + ) + self.hidden_size = hidden_size + self.cross_attention_dim = cross_attention_dim + self.rank = rank + q_rank = kwargs.pop("q_rank", None) + q_hidden_size = kwargs.pop("q_hidden_size", None) + q_rank = q_rank if q_rank is not None else rank + q_hidden_size = q_hidden_size if q_hidden_size is not None else hidden_size + v_rank = kwargs.pop("v_rank", None) + v_hidden_size = kwargs.pop("v_hidden_size", None) + v_rank = v_rank if v_rank is not None else rank + v_hidden_size = v_hidden_size if v_hidden_size is not None else hidden_size + out_rank = kwargs.pop("out_rank", None) + out_hidden_size = kwargs.pop("out_hidden_size", None) + out_rank = out_rank if out_rank is not None else rank + out_hidden_size = ( + out_hidden_size if out_hidden_size is not None else hidden_size + ) + self.to_q_lora = LoRALinearLayer( + q_hidden_size, q_hidden_size, q_rank, network_alpha + ) + self.to_k_lora = LoRALinearLayer( + cross_attention_dim or hidden_size, hidden_size, rank, network_alpha + ) + self.to_v_lora = LoRALinearLayer( + cross_attention_dim or v_hidden_size, v_hidden_size, v_rank, network_alpha + ) + self.to_out_lora = LoRALinearLayer( + out_hidden_size, out_hidden_size, out_rank, network_alpha + ) + + def __call__( + self, attn: Attention, hidden_states: paddle.Tensor, **kwargs + ) -> paddle.Tensor: + self_cls_name = self.__class__.__name__ + deprecate( + self_cls_name, + "0.26.0", + f"Make sure use {self_cls_name[4:]} instead by settingLoRA layers to `self.{{to_q,to_k,to_v,to_out[0]}}.lora_layer` respectively. This will be done automatically when using `LoraLoaderMixin.load_lora_weights`", + ) + attn.to_q.lora_layer = self.to_q_lora.to(hidden_states.place) + attn.to_k.lora_layer = self.to_k_lora.to(hidden_states.place) + attn.to_v.lora_layer = self.to_v_lora.to(hidden_states.place) + attn.to_out[0].lora_layer = self.to_out_lora.to(hidden_states.place) + attn._modules.pop("processor") + attn.processor = AttnProcessor2_0() + return attn.processor(attn, hidden_states, **kwargs) + + +class LoRAXFormersAttnProcessor(paddle.nn.Layer): + """ + Processor for implementing the LoRA attention mechanism with memory efficient attention using xFormers. + + Args: + hidden_size (`int`, *optional*): + The hidden size of the attention layer. + cross_attention_dim (`int`, *optional*): + The number of channels in the `encoder_hidden_states`. + rank (`int`, defaults to 4): + The dimension of the LoRA update matrices. + attention_op (`Callable`, *optional*, defaults to `None`): + The base + [operator](https://facebookresearch.github.io/xformers/components/ops.html#xformers.ops.AttentionOpBase) to + use as the attention operator. It is recommended to set to `None`, and allow xFormers to choose the best + operator. + network_alpha (`int`, *optional*): + Equivalent to `alpha` but it's usage is specific to Kohya (A1111) style LoRAs. + kwargs (`dict`): + Additional keyword arguments to pass to the `LoRALinearLayer` layers. + """ + + def __init__( + self, + hidden_size: int, + cross_attention_dim: int, + rank: int = 4, + attention_op: Optional[Callable] = None, + network_alpha: Optional[int] = None, + **kwargs, + ): + super().__init__() + self.hidden_size = hidden_size + self.cross_attention_dim = cross_attention_dim + self.rank = rank + self.attention_op = attention_op + q_rank = kwargs.pop("q_rank", None) + q_hidden_size = kwargs.pop("q_hidden_size", None) + q_rank = q_rank if q_rank is not None else rank + q_hidden_size = q_hidden_size if q_hidden_size is not None else hidden_size + v_rank = kwargs.pop("v_rank", None) + v_hidden_size = kwargs.pop("v_hidden_size", None) + v_rank = v_rank if v_rank is not None else rank + v_hidden_size = v_hidden_size if v_hidden_size is not None else hidden_size + out_rank = kwargs.pop("out_rank", None) + out_hidden_size = kwargs.pop("out_hidden_size", None) + out_rank = out_rank if out_rank is not None else rank + out_hidden_size = ( + out_hidden_size if out_hidden_size is not None else hidden_size + ) + self.to_q_lora = LoRALinearLayer( + q_hidden_size, q_hidden_size, q_rank, network_alpha + ) + self.to_k_lora = LoRALinearLayer( + cross_attention_dim or hidden_size, hidden_size, rank, network_alpha + ) + self.to_v_lora = LoRALinearLayer( + cross_attention_dim or v_hidden_size, v_hidden_size, v_rank, network_alpha + ) + self.to_out_lora = LoRALinearLayer( + out_hidden_size, out_hidden_size, out_rank, network_alpha + ) + + def __call__( + self, attn: Attention, hidden_states: paddle.Tensor, **kwargs + ) -> paddle.Tensor: + self_cls_name = self.__class__.__name__ + deprecate( + self_cls_name, + "0.26.0", + f"Make sure use {self_cls_name[4:]} instead by settingLoRA layers to `self.{{to_q,to_k,to_v,add_k_proj,add_v_proj,to_out[0]}}.lora_layer` respectively. This will be done automatically when using `LoraLoaderMixin.load_lora_weights`", + ) + attn.to_q.lora_layer = self.to_q_lora.to(hidden_states.place) + attn.to_k.lora_layer = self.to_k_lora.to(hidden_states.place) + attn.to_v.lora_layer = self.to_v_lora.to(hidden_states.place) + attn.to_out[0].lora_layer = self.to_out_lora.to(hidden_states.place) + attn._modules.pop("processor") + attn.processor = XFormersAttnProcessor() + return attn.processor(attn, hidden_states, **kwargs) + + +class LoRAAttnAddedKVProcessor(paddle.nn.Layer): + """ + Processor for implementing the LoRA attention mechanism with extra learnable key and value matrices for the text + encoder. + + Args: + hidden_size (`int`, *optional*): + The hidden size of the attention layer. + cross_attention_dim (`int`, *optional*, defaults to `None`): + The number of channels in the `encoder_hidden_states`. + rank (`int`, defaults to 4): + The dimension of the LoRA update matrices. + network_alpha (`int`, *optional*): + Equivalent to `alpha` but it's usage is specific to Kohya (A1111) style LoRAs. + kwargs (`dict`): + Additional keyword arguments to pass to the `LoRALinearLayer` layers. + """ + + def __init__( + self, + hidden_size: int, + cross_attention_dim: Optional[int] = None, + rank: int = 4, + network_alpha: Optional[int] = None, + ): + super().__init__() + self.hidden_size = hidden_size + self.cross_attention_dim = cross_attention_dim + self.rank = rank + self.to_q_lora = LoRALinearLayer(hidden_size, hidden_size, rank, network_alpha) + self.add_k_proj_lora = LoRALinearLayer( + cross_attention_dim or hidden_size, hidden_size, rank, network_alpha + ) + self.add_v_proj_lora = LoRALinearLayer( + cross_attention_dim or hidden_size, hidden_size, rank, network_alpha + ) + self.to_k_lora = LoRALinearLayer(hidden_size, hidden_size, rank, network_alpha) + self.to_v_lora = LoRALinearLayer(hidden_size, hidden_size, rank, network_alpha) + self.to_out_lora = LoRALinearLayer( + hidden_size, hidden_size, rank, network_alpha + ) + + def __call__( + self, attn: Attention, hidden_states: paddle.Tensor, **kwargs + ) -> paddle.Tensor: + self_cls_name = self.__class__.__name__ + deprecate( + self_cls_name, + "0.26.0", + f"Make sure use {self_cls_name[4:]} instead by settingLoRA layers to `self.{{to_q,to_k,to_v,add_k_proj,add_v_proj,to_out[0]}}.lora_layer` respectively. This will be done automatically when using `LoraLoaderMixin.load_lora_weights`", + ) + attn.to_q.lora_layer = self.to_q_lora.to(hidden_states.place) + attn.to_k.lora_layer = self.to_k_lora.to(hidden_states.place) + attn.to_v.lora_layer = self.to_v_lora.to(hidden_states.place) + attn.to_out[0].lora_layer = self.to_out_lora.to(hidden_states.place) + attn._modules.pop("processor") + attn.processor = AttnAddedKVProcessor() + return attn.processor(attn, hidden_states, **kwargs) + + +class IPAdapterAttnProcessor(paddle.nn.Layer): + """ + Attention processor for Multiple IP-Adapters. + + Args: + hidden_size (`int`): + The hidden size of the attention layer. + cross_attention_dim (`int`): + The number of channels in the `encoder_hidden_states`. + num_tokens (`int`, `Tuple[int]` or `List[int]`, defaults to `(4,)`): + The context length of the image features. + scale (`float` or List[`float`], defaults to 1.0): + the weight scale of image prompt. + """ + + def __init__( + self, hidden_size, cross_attention_dim=None, num_tokens=(4,), scale=1.0 + ): + super().__init__() + self.hidden_size = hidden_size + self.cross_attention_dim = cross_attention_dim + if not isinstance(num_tokens, (tuple, list)): + num_tokens = [num_tokens] + self.num_tokens = num_tokens + if not isinstance(scale, list): + scale = [scale] * len(num_tokens) + if len(scale) != len(num_tokens): + raise ValueError( + "`scale` should be a list of integers with the same length as `num_tokens`." + ) + self.scale = scale + self.to_k_ip = paddle.nn.LayerList( + sublayers=[ + paddle.nn.Linear( + in_features=cross_attention_dim, + out_features=hidden_size, + bias_attr=False, + ) + for _ in range(len(num_tokens)) + ] + ) + self.to_v_ip = paddle.nn.LayerList( + sublayers=[ + paddle.nn.Linear( + in_features=cross_attention_dim, + out_features=hidden_size, + bias_attr=False, + ) + for _ in range(len(num_tokens)) + ] + ) + + def __call__( + self, + attn: Attention, + hidden_states: paddle.Tensor, + encoder_hidden_states: Optional[paddle.Tensor] = None, + attention_mask: Optional[paddle.Tensor] = None, + temb: Optional[paddle.Tensor] = None, + scale: float = 1.0, + ip_adapter_masks: Optional[paddle.Tensor] = None, + ): + residual = hidden_states + if encoder_hidden_states is not None: + if isinstance(encoder_hidden_states, tuple): + encoder_hidden_states, ip_hidden_states = encoder_hidden_states + else: + deprecation_message = "You have passed a tensor as `encoder_hidden_states`. This is deprecated and will be removed in a future release. Please make sure to update your script to pass `encoder_hidden_states` as a tuple to suppress this warning." + deprecate( + "encoder_hidden_states not a tuple", + "1.0.0", + deprecation_message, + standard_warn=False, + ) + end_pos = encoder_hidden_states.shape[1] - self.num_tokens[0] + encoder_hidden_states, ip_hidden_states = ( + encoder_hidden_states[:, :end_pos, :], + [encoder_hidden_states[:, end_pos:, :]], + ) + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + input_ndim = hidden_states.ndim + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view( + batch_size, channel, height * width + ).transpose(1, 2) + batch_size, sequence_length, _ = ( + hidden_states.shape + if encoder_hidden_states is None + else encoder_hidden_states.shape + ) + attention_mask = attn.prepare_attention_mask( + attention_mask, sequence_length, batch_size + ) + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose( + 1, 2 + ) + query = attn.to_q(hidden_states) + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states( + encoder_hidden_states + ) + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + query = attn.head_to_batch_dim(query) + key = attn.head_to_batch_dim(key) + value = attn.head_to_batch_dim(value) + attention_probs = attn.get_attention_scores(query, key, attention_mask) + hidden_states = paddle.bmm(attention_probs, value) + hidden_states = attn.batch_to_head_dim(hidden_states) + if ip_adapter_masks is not None: + if not isinstance(ip_adapter_masks, List): + ip_adapter_masks = list(ip_adapter_masks.unsqueeze(1)) + if not len(ip_adapter_masks) == len(self.scale) == len(ip_hidden_states): + raise ValueError( + f"Length of ip_adapter_masks array ({len(ip_adapter_masks)}) must match length of self.scale array ({len(self.scale)}) and number of ip_hidden_states ({len(ip_hidden_states)})" + ) + else: + for index, (mask, scale, ip_state) in enumerate( + zip(ip_adapter_masks, self.scale, ip_hidden_states) + ): + if not isinstance(mask, paddle.Tensor) or mask.ndim != 4: + raise ValueError( + "Each element of the ip_adapter_masks array should be a tensor with shape [1, num_images_for_ip_adapter, height, width]. Please use `IPAdapterMaskProcessor` to preprocess your mask" + ) + if mask.shape[1] != ip_state.shape[1]: + raise ValueError( + f"Number of masks ({mask.shape[1]}) does not match number of ip images ({ip_state.shape[1]}) at index {index}" + ) + if isinstance(scale, list) and not len(scale) == mask.shape[1]: + raise ValueError( + f"Number of masks ({mask.shape[1]}) does not match number of scales ({len(scale)}) at index {index}" + ) + else: + ip_adapter_masks = [None] * len(self.scale) + for current_ip_hidden_states, scale, to_k_ip, to_v_ip, mask in zip( + ip_hidden_states, self.scale, self.to_k_ip, self.to_v_ip, ip_adapter_masks + ): + skip = False + if isinstance(scale, list): + if all(s == 0 for s in scale): + skip = True + elif scale == 0: + skip = True + if not skip: + if mask is not None: + if not isinstance(scale, list): + scale = [scale] * mask.shape[1] + current_num_images = mask.shape[1] + for i in range(current_num_images): + ip_key = to_k_ip(current_ip_hidden_states[:, i, :, :]) + ip_value = to_v_ip(current_ip_hidden_states[:, i, :, :]) + ip_key = attn.head_to_batch_dim(ip_key) + ip_value = attn.head_to_batch_dim(ip_value) + ip_attention_probs = attn.get_attention_scores( + query, ip_key, None + ) + _current_ip_hidden_states = paddle.bmm( + ip_attention_probs, ip_value + ) + _current_ip_hidden_states = attn.batch_to_head_dim( + _current_ip_hidden_states + ) + mask_downsample = IPAdapterMaskProcessor.downsample( + mask[:, i, :, :], + batch_size, + _current_ip_hidden_states.shape[1], + _current_ip_hidden_states.shape[2], + ) + mask_downsample = mask_downsample.to( + dtype=query.dtype, device=query.place + ) + hidden_states = hidden_states + scale[i] * ( + _current_ip_hidden_states * mask_downsample + ) + else: + ip_key = to_k_ip(current_ip_hidden_states) + ip_value = to_v_ip(current_ip_hidden_states) + ip_key = attn.head_to_batch_dim(ip_key) + ip_value = attn.head_to_batch_dim(ip_value) + ip_attention_probs = attn.get_attention_scores(query, ip_key, None) + current_ip_hidden_states = paddle.bmm(ip_attention_probs, ip_value) + current_ip_hidden_states = attn.batch_to_head_dim( + current_ip_hidden_states + ) + hidden_states = hidden_states + scale * current_ip_hidden_states + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape( + batch_size, channel, height, width + ) + if attn.residual_connection: + hidden_states = hidden_states + residual + hidden_states = hidden_states / attn.rescale_output_factor + return hidden_states + + +class IPAdapterAttnProcessor2_0(paddle.nn.Layer): + """ + Attention processor for IP-Adapter for PyTorch 2.0. + + Args: + hidden_size (`int`): + The hidden size of the attention layer. + cross_attention_dim (`int`): + The number of channels in the `encoder_hidden_states`. + num_tokens (`int`, `Tuple[int]` or `List[int]`, defaults to `(4,)`): + The context length of the image features. + scale (`float` or `List[float]`, defaults to 1.0): + the weight scale of image prompt. + """ + + def __init__( + self, hidden_size, cross_attention_dim=None, num_tokens=(4,), scale=1.0 + ): + super().__init__() +>>>>>> if not hasattr(torch.nn.functional, "scaled_dot_product_attention"): + raise ImportError( + f"{self.__class__.__name__} requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0." + ) + self.hidden_size = hidden_size + self.cross_attention_dim = cross_attention_dim + if not isinstance(num_tokens, (tuple, list)): + num_tokens = [num_tokens] + self.num_tokens = num_tokens + if not isinstance(scale, list): + scale = [scale] * len(num_tokens) + if len(scale) != len(num_tokens): + raise ValueError( + "`scale` should be a list of integers with the same length as `num_tokens`." + ) + self.scale = scale + self.to_k_ip = paddle.nn.LayerList( + sublayers=[ + paddle.nn.Linear( + in_features=cross_attention_dim, + out_features=hidden_size, + bias_attr=False, + ) + for _ in range(len(num_tokens)) + ] + ) + self.to_v_ip = paddle.nn.LayerList( + sublayers=[ + paddle.nn.Linear( + in_features=cross_attention_dim, + out_features=hidden_size, + bias_attr=False, + ) + for _ in range(len(num_tokens)) + ] + ) + + def __call__( + self, + attn: Attention, + hidden_states: paddle.Tensor, + encoder_hidden_states: Optional[paddle.Tensor] = None, + attention_mask: Optional[paddle.Tensor] = None, + temb: Optional[paddle.Tensor] = None, + scale: float = 1.0, + ip_adapter_masks: Optional[paddle.Tensor] = None, + ): + residual = hidden_states + if encoder_hidden_states is not None: + if isinstance(encoder_hidden_states, tuple): + encoder_hidden_states, ip_hidden_states = encoder_hidden_states + else: + deprecation_message = "You have passed a tensor as `encoder_hidden_states`. This is deprecated and will be removed in a future release. Please make sure to update your script to pass `encoder_hidden_states` as a tuple to suppress this warning." + deprecate( + "encoder_hidden_states not a tuple", + "1.0.0", + deprecation_message, + standard_warn=False, + ) + end_pos = encoder_hidden_states.shape[1] - self.num_tokens[0] + encoder_hidden_states, ip_hidden_states = ( + encoder_hidden_states[:, :end_pos, :], + [encoder_hidden_states[:, end_pos:, :]], + ) + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + input_ndim = hidden_states.ndim + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view( + batch_size, channel, height * width + ).transpose(1, 2) + batch_size, sequence_length, _ = ( + hidden_states.shape + if encoder_hidden_states is None + else encoder_hidden_states.shape + ) + if attention_mask is not None: + attention_mask = attn.prepare_attention_mask( + attention_mask, sequence_length, batch_size + ) + attention_mask = attention_mask.view( + batch_size, attn.heads, -1, attention_mask.shape[-1] + ) + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose( + 1, 2 + ) + query = attn.to_q(hidden_states) + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states( + encoder_hidden_states + ) + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + inner_dim = key.shape[-1] + head_dim = inner_dim // attn.heads + query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + hidden_states = paddle.nn.functional.scaled_dot_product_attention( + query.transpose([0, 2, 1, 3]), + key.transpose([0, 2, 1, 3]), + value.transpose([0, 2, 1, 3]), + attn_mask=attention_mask, + dropout_p=0.0, + is_causal=False, + ).transpose([0, 2, 1, 3]) + hidden_states = hidden_states.transpose(1, 2).reshape( + batch_size, -1, attn.heads * head_dim + ) + hidden_states = hidden_states.to(query.dtype) + if ip_adapter_masks is not None: + if not isinstance(ip_adapter_masks, List): + ip_adapter_masks = list(ip_adapter_masks.unsqueeze(1)) + if not len(ip_adapter_masks) == len(self.scale) == len(ip_hidden_states): + raise ValueError( + f"Length of ip_adapter_masks array ({len(ip_adapter_masks)}) must match length of self.scale array ({len(self.scale)}) and number of ip_hidden_states ({len(ip_hidden_states)})" + ) + else: + for index, (mask, scale, ip_state) in enumerate( + zip(ip_adapter_masks, self.scale, ip_hidden_states) + ): + if not isinstance(mask, paddle.Tensor) or mask.ndim != 4: + raise ValueError( + "Each element of the ip_adapter_masks array should be a tensor with shape [1, num_images_for_ip_adapter, height, width]. Please use `IPAdapterMaskProcessor` to preprocess your mask" + ) + if mask.shape[1] != ip_state.shape[1]: + raise ValueError( + f"Number of masks ({mask.shape[1]}) does not match number of ip images ({ip_state.shape[1]}) at index {index}" + ) + if isinstance(scale, list) and not len(scale) == mask.shape[1]: + raise ValueError( + f"Number of masks ({mask.shape[1]}) does not match number of scales ({len(scale)}) at index {index}" + ) + else: + ip_adapter_masks = [None] * len(self.scale) + for current_ip_hidden_states, scale, to_k_ip, to_v_ip, mask in zip( + ip_hidden_states, self.scale, self.to_k_ip, self.to_v_ip, ip_adapter_masks + ): + skip = False + if isinstance(scale, list): + if all(s == 0 for s in scale): + skip = True + elif scale == 0: + skip = True + if not skip: + if mask is not None: + if not isinstance(scale, list): + scale = [scale] * mask.shape[1] + current_num_images = mask.shape[1] + for i in range(current_num_images): + ip_key = to_k_ip(current_ip_hidden_states[:, i, :, :]) + ip_value = to_v_ip(current_ip_hidden_states[:, i, :, :]) + ip_key = ip_key.view( + batch_size, -1, attn.heads, head_dim + ).transpose(1, 2) + ip_value = ip_value.view( + batch_size, -1, attn.heads, head_dim + ).transpose(1, 2) + _current_ip_hidden_states = ( + paddle.nn.functional.scaled_dot_product_attention( + query.transpose([0, 2, 1, 3]), + ip_key.transpose([0, 2, 1, 3]), + ip_value.transpose([0, 2, 1, 3]), + attn_mask=None, + dropout_p=0.0, + is_causal=False, + ).transpose([0, 2, 1, 3]) + ) + _current_ip_hidden_states = _current_ip_hidden_states.transpose( + 1, 2 + ).reshape(batch_size, -1, attn.heads * head_dim) + _current_ip_hidden_states = _current_ip_hidden_states.to( + query.dtype + ) + mask_downsample = IPAdapterMaskProcessor.downsample( + mask[:, i, :, :], + batch_size, + _current_ip_hidden_states.shape[1], + _current_ip_hidden_states.shape[2], + ) + mask_downsample = mask_downsample.to( + dtype=query.dtype, device=query.place + ) + hidden_states = hidden_states + scale[i] * ( + _current_ip_hidden_states * mask_downsample + ) + else: + ip_key = to_k_ip(current_ip_hidden_states) + ip_value = to_v_ip(current_ip_hidden_states) + ip_key = ip_key.view( + batch_size, -1, attn.heads, head_dim + ).transpose(1, 2) + ip_value = ip_value.view( + batch_size, -1, attn.heads, head_dim + ).transpose(1, 2) + current_ip_hidden_states = ( + paddle.nn.functional.scaled_dot_product_attention( + query.transpose([0, 2, 1, 3]), + ip_key.transpose([0, 2, 1, 3]), + ip_value.transpose([0, 2, 1, 3]), + attn_mask=None, + dropout_p=0.0, + is_causal=False, + ).transpose([0, 2, 1, 3]) + ) + current_ip_hidden_states = current_ip_hidden_states.transpose( + 1, 2 + ).reshape(batch_size, -1, attn.heads * head_dim) + current_ip_hidden_states = current_ip_hidden_states.to(query.dtype) + hidden_states = hidden_states + scale * current_ip_hidden_states + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape( + batch_size, channel, height, width + ) + if attn.residual_connection: + hidden_states = hidden_states + residual + hidden_states = hidden_states / attn.rescale_output_factor + return hidden_states + + +LORA_ATTENTION_PROCESSORS = ( + LoRAAttnProcessor, + LoRAAttnProcessor2_0, + LoRAXFormersAttnProcessor, + LoRAAttnAddedKVProcessor, +) +ADDED_KV_ATTENTION_PROCESSORS = ( + AttnAddedKVProcessor, + SlicedAttnAddedKVProcessor, + AttnAddedKVProcessor2_0, + XFormersAttnAddedKVProcessor, + LoRAAttnAddedKVProcessor, +) +CROSS_ATTENTION_PROCESSORS = ( + AttnProcessor, + AttnProcessor2_0, + XFormersAttnProcessor, + SlicedAttnProcessor, + LoRAAttnProcessor, + LoRAAttnProcessor2_0, + LoRAXFormersAttnProcessor, + IPAdapterAttnProcessor, + IPAdapterAttnProcessor2_0, +) +AttentionProcessor = Union[ + AttnProcessor, + AttnProcessor2_0, + FusedAttnProcessor2_0, + XFormersAttnProcessor, + SlicedAttnProcessor, + AttnAddedKVProcessor, + SlicedAttnAddedKVProcessor, + AttnAddedKVProcessor2_0, + XFormersAttnAddedKVProcessor, + CustomDiffusionAttnProcessor, + CustomDiffusionXFormersAttnProcessor, + CustomDiffusionAttnProcessor2_0, + LoRAAttnProcessor, + LoRAAttnProcessor2_0, + LoRAXFormersAttnProcessor, + LoRAAttnAddedKVProcessor, +] diff --git a/paddlespeech/t2s/modules/flow/convolution.py b/paddlespeech/t2s/modules/flow/convolution.py new file mode 100644 index 000000000..9f38479d8 --- /dev/null +++ b/paddlespeech/t2s/modules/flow/convolution.py @@ -0,0 +1,99 @@ +import paddle + +"""ConvolutionModule definition.""" +from typing import Tuple + + +class ConvolutionModule(paddle.nn.Layer): + """ConvolutionModule in Conformer model.""" + + def __init__( + self, + channels: int, + kernel_size: int = 15, + activation: paddle.nn.Layer = paddle.nn.ReLU(), + norm: str = "batch_norm", + causal: bool = False, + bias: bool = True, + ): + """Construct an ConvolutionModule object. + Args: + channels (int): The number of channels of conv layers. + kernel_size (int): Kernel size of conv layers. + causal (int): Whether use causal convolution or not + """ + super().__init__() + self.pointwise_conv1 = paddle.nn.Conv1d( + channels, 2 * channels, kernel_size=1, stride=1, padding=0, bias=bias + ) + if causal: + padding = 0 + self.lorder = kernel_size - 1 + else: + assert (kernel_size - 1) % 2 == 0 + padding = (kernel_size - 1) // 2 + self.lorder = 0 + self.depthwise_conv = paddle.nn.Conv1d( + channels, + channels, + kernel_size, + stride=1, + padding=padding, + groups=channels, + bias=bias, + ) + assert norm in ["batch_norm", "layer_norm"] + if norm == "batch_norm": + self.use_layer_norm = False + self.norm = paddle.nn.BatchNorm1D(num_features=channels) + else: + self.use_layer_norm = True + self.norm = paddle.nn.LayerNorm(normalized_shape=channels) + self.pointwise_conv2 = paddle.nn.Conv1d( + channels, channels, kernel_size=1, stride=1, padding=0, bias=bias + ) + self.activation = activation + + def forward( + self, + x: paddle.Tensor, + mask_pad: paddle.Tensor = paddle.ones((0, 0, 0), dtype=paddle.bool), + cache: paddle.Tensor = paddle.zeros((0, 0, 0)), + ) -> Tuple[paddle.Tensor, paddle.Tensor]: + """Compute convolution module. + Args: + x (torch.Tensor): Input tensor (#batch, time, channels). + mask_pad (torch.Tensor): used for batch padding (#batch, 1, time), + (0, 0, 0) means fake mask. + cache (torch.Tensor): left context cache, it is only + used in causal convolution (#batch, channels, cache_t), + (0, 0, 0) meas fake cache. + Returns: + torch.Tensor: Output tensor (#batch, time, channels). + """ + x = x.transpose(1, 2) + if mask_pad.size(2) > 0: + x.masked_fill_(~mask_pad, 0.0) + if self.lorder > 0: + if cache.size(2) == 0: + x = paddle.compat.pad(x, (self.lorder, 0), "constant", 0.0) + else: + assert cache.size(0) == x.size(0) + assert cache.size(1) == x.size(1) + x = paddle.cat((cache, x), dim=2) + assert x.size(2) > self.lorder + new_cache = x[:, :, -self.lorder :] + else: + new_cache = paddle.zeros((0, 0, 0), dtype=x.dtype, device=x.place) + x = self.pointwise_conv1(x) + x = paddle.nn.functional.glu(x=x, axis=1) + x = self.depthwise_conv(x) + if self.use_layer_norm: + x = x.transpose(1, 2) + x = self.activation(self.norm(x)) + if self.use_layer_norm: + x = x.transpose(1, 2) + x = self.pointwise_conv2(x) + if mask_pad.size(2) > 0: + x.masked_fill_(~mask_pad, 0.0) + return x.transpose(1, 2), new_cache diff --git a/paddlespeech/t2s/modules/flow/decoder.py b/paddlespeech/t2s/modules/flow/decoder.py index 4c5208b50..7744656b1 100644 --- a/paddlespeech/t2s/modules/flow/decoder.py +++ b/paddlespeech/t2s/modules/flow/decoder.py @@ -11,15 +11,15 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -from typing import Tuple +from typing import Tuple, Any, Dict, Optional import paddle +import math from paddle import nn import paddle.nn.functional as F from einops import pack, rearrange, repeat -from cosyvoice.utils.common import mask_to_bias -from cosyvoice.utils.mask import add_optional_chunk_mask -from matcha.models.components.decoder import SinusoidalPosEmb, Block1D, ResnetBlock1D, Downsample1D, TimestepEmbedding, Upsample1D -from .attention import BasicTransformerBlock +from paddlespeech.t2s.models.CosyVoice.common import mask_to_bias +from paddlespeech.t2s.models.CosyVoice.mask import add_optional_chunk_mask +from .matcha_transformer import BasicTransformerBlock def get_activation(act_fn): if act_fn == "silu": @@ -110,14 +110,15 @@ class TimestepEmbedding(nn.Layer): if condition is not None and self.cond_proj is not None: sample = sample + self.cond_proj(condition) sample = self.linear_1(sample) - + # print("sample2:",sample) if self.act is not None: sample = self.act(sample) - + # print("sample3:",sample) sample = self.linear_2(sample) - + # print("sample4:",sample) if self.post_act is not None: sample = self.post_act(sample) + # print("sample5:",sample) return sample class Upsample1D(nn.Layer): @@ -160,17 +161,17 @@ class Upsample1D(nn.Layer): return outputs -class Transpose(nn.Module): +class Transpose(nn.Layer): def __init__(self, dim0: int, dim1: int): super().__init__() self.dim0 = dim0 self.dim1 = dim1 def forward(self, x: paddle.Tensor) -> paddle.Tensor: - x = paddle.transpose(x, (self.dim0, self.dim1)) + x = paddle.transpose(x, [0, self.dim1, self.dim0]) return x -class CausalConv1d(nn.Conv1d): +class CausalConv1d(nn.Conv1D): def __init__( self, in_channels: int, @@ -332,8 +333,7 @@ def add_optional_chunk_mask(xs: paddle.Tensor, chunk_masks = masks & chunk_masks # (B, L, L) else: chunk_masks = masks - - assert chunk_masks.dtype == 'bool' + assert chunk_masks.dtype == paddle.bool if (chunk_masks.sum(axis=-1) == 0).sum().item() != 0: print('get chunk_masks all false at some timestep, force set to true, make sure they are masked in future computation!') all_false_mask = chunk_masks.sum(axis=-1) == 0 @@ -342,8 +342,8 @@ def add_optional_chunk_mask(xs: paddle.Tensor, return chunk_masks def mask_to_bias(mask: paddle.Tensor, dtype: str) -> paddle.Tensor: - assert mask.dtype == 'bool', "Input mask must be of boolean type" - assert dtype in ['float32', 'bfloat16', 'float16'], f"Unsupported dtype: {dtype}" + assert mask.dtype == paddle.bool, "Input mask must be of boolean type" + assert dtype in [paddle.float32, paddle.bfloat16, paddle.float16], f"Unsupported dtype: {dtype}" mask = mask.astype(dtype) mask = (1.0 - mask) * -1.0e+10 @@ -489,7 +489,6 @@ class ConditionalDecoder(nn.Layer): t = self.time_mlp(t) x = pack([x, mu], "b * t")[0] - if spks is not None: spks = repeat(spks, "b c -> b c t", t=x.shape[-1]) x = pack([x, spks], "b * t")[0] @@ -667,15 +666,16 @@ class CausalConditionalDecoder(nn.Layer): if isinstance(m, nn.Conv1D): nn.initializer.KaimingNormal(m.weight, nonlinearity='relu') if m.bias is not None: - nn.initializer.Constant(m.bias, value=0) + initializer = nn.initializer.Constant(value=0) + initializer(m.bias) elif isinstance(m, nn.GroupNorm): nn.initializer.Constant(m.weight, value=1) nn.initializer.Constant(m.bias, value=0) elif isinstance(m, nn.Linear): nn.initializer.KaimingNormal(m.weight, nonlinearity='relu') if m.bias is not None: - nn.initializer.Constant(m.bias, value=0) - + initializer = nn.initializer.Constant(value=0) + initializer(m.bias) def forward(self, x, mask, mu, t, spks=None, cond=None, streaming=False): """Forward pass of the UNet1DConditional model. @@ -693,9 +693,7 @@ class CausalConditionalDecoder(nn.Layer): """ t = self.time_embeddings(t).astype(t.dtype) # 使用 astype 代替 .to(t.dtype) t = self.time_mlp(t) - x = pack([x, mu], "b * t")[0] # 假设 pack 函数已实现 - if spks is not None: spks = repeat(spks, "b c -> b c t", t=x.shape[-1]) # 假设 repeat 函数已实现 x = pack([x, spks], "b * t")[0] diff --git a/paddlespeech/t2s/modules/flow/diffusers_activatioins.py b/paddlespeech/t2s/modules/flow/diffusers_activatioins.py new file mode 100644 index 000000000..031930167 --- /dev/null +++ b/paddlespeech/t2s/modules/flow/diffusers_activatioins.py @@ -0,0 +1,132 @@ +import paddle + +from ..utils import deprecate +from ..utils.import_utils import is_torch_npu_available + +if is_torch_npu_available(): + import torch_npu +ACTIVATION_FUNCTIONS = { + "swish": paddle.nn.SiLU(), + "silu": paddle.nn.SiLU(), + "mish": paddle.nn.Mish(), + "gelu": paddle.nn.GELU(), + "relu": paddle.nn.ReLU(), +} + + +def get_activation(act_fn: str) -> paddle.nn.Layer: + """Helper function to get activation function from string. + + Args: + act_fn (str): Name of activation function. + + Returns: + nn.Module: Activation function. + """ + act_fn = act_fn.lower() + if act_fn in ACTIVATION_FUNCTIONS: + return ACTIVATION_FUNCTIONS[act_fn] + else: + raise ValueError(f"Unsupported activation function: {act_fn}") + + +class FP32SiLU(paddle.nn.Layer): + """ + SiLU activation function with input upcasted to torch.float32. + """ + + def __init__(self): + super().__init__() + + def forward(self, inputs: paddle.Tensor) -> paddle.Tensor: + return paddle.nn.functional.silu(inputs.float(), inplace=False).to(inputs.dtype) + + +class GELU(paddle.nn.Layer): + """ + GELU activation function with tanh approximation support with `approximate="tanh"`. + + Parameters: + dim_in (`int`): The number of channels in the input. + dim_out (`int`): The number of channels in the output. + approximate (`str`, *optional*, defaults to `"none"`): If `"tanh"`, use tanh approximation. + bias (`bool`, defaults to True): Whether to use a bias in the linear layer. + """ + + def __init__( + self, dim_in: int, dim_out: int, approximate: str = "none", bias: bool = True + ): + super().__init__() + self.proj = paddle.nn.Linear( + in_features=dim_in, out_features=dim_out, bias_attr=bias + ) + self.approximate = approximate + + def gelu(self, gate: paddle.Tensor) -> paddle.Tensor: + if gate.device.type != "mps": + return paddle.nn.functional.gelu(gate, approximate=self.approximate) + return paddle.nn.functional.gelu( + gate.to(dtype=paddle.float32), approximate=self.approximate + ).to(dtype=gate.dtype) + + def forward(self, hidden_states): + hidden_states = self.proj(hidden_states) + hidden_states = self.gelu(hidden_states) + return hidden_states + + +class GEGLU(paddle.nn.Layer): + """ + A [variant](https://arxiv.org/abs/2002.05202) of the gated linear unit activation function. + + Parameters: + dim_in (`int`): The number of channels in the input. + dim_out (`int`): The number of channels in the output. + bias (`bool`, defaults to True): Whether to use a bias in the linear layer. + """ + + def __init__(self, dim_in: int, dim_out: int, bias: bool = True): + super().__init__() + self.proj = paddle.nn.Linear( + in_features=dim_in, out_features=dim_out * 2, bias_attr=bias + ) + + def gelu(self, gate: paddle.Tensor) -> paddle.Tensor: + if gate.device.type != "mps": + return paddle.nn.functional.gelu(gate) + return paddle.nn.functional.gelu(gate.to(dtype=paddle.float32)).to( + dtype=gate.dtype + ) + + def forward(self, hidden_states, *args, **kwargs): + if len(args) > 0 or kwargs.get("scale", None) is not None: + deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." + deprecate("scale", "1.0.0", deprecation_message) + hidden_states = self.proj(hidden_states) + if is_torch_npu_available(): + return torch_npu.npu_geglu(hidden_states, dim=-1, approximate=1)[0] + else: + hidden_states, gate = hidden_states.chunk(2, dim=-1) + return hidden_states * self.gelu(gate) + + +class ApproximateGELU(paddle.nn.Layer): + """ + The approximate form of the Gaussian Error Linear Unit (GELU). For more details, see section 2 of this + [paper](https://arxiv.org/abs/1606.08415). + + Parameters: + dim_in (`int`): The number of channels in the input. + dim_out (`int`): The number of channels in the output. + bias (`bool`, defaults to True): Whether to use a bias in the linear layer. + """ + + def __init__(self, dim_in: int, dim_out: int, bias: bool = True): + super().__init__() + self.proj = paddle.nn.Linear( + in_features=dim_in, out_features=dim_out, bias_attr=bias + ) + + def forward(self, x: paddle.Tensor) -> paddle.Tensor: + x = self.proj(x) + return x * paddle.nn.functional.sigmoid(1.702 * x) diff --git a/paddlespeech/t2s/modules/flow/encoder_layer.py b/paddlespeech/t2s/modules/flow/encoder_layer.py new file mode 100644 index 000000000..d7debadf9 --- /dev/null +++ b/paddlespeech/t2s/modules/flow/encoder_layer.py @@ -0,0 +1,205 @@ +import paddle + +"""Encoder self-attention layer definition.""" +from typing import Optional, Tuple + + +class TransformerEncoderLayer(paddle.nn.Layer): + """Encoder layer module. + + Args: + size (int): Input dimension. + self_attn (torch.nn.Module): Self-attention module instance. + `MultiHeadedAttention` or `RelPositionMultiHeadedAttention` + instance can be used as the argument. + feed_forward (torch.nn.Module): Feed-forward module instance. + `PositionwiseFeedForward`, instance can be used as the argument. + dropout_rate (float): Dropout rate. + normalize_before (bool): + True: use layer_norm before each sub-block. + False: to use layer_norm after each sub-block. + """ + + def __init__( + self, + size: int, + self_attn: paddle.nn.Layer, + feed_forward: paddle.nn.Layer, + dropout_rate: float, + normalize_before: bool = True, + ): + """Construct an EncoderLayer object.""" + super().__init__() + self.self_attn = self_attn + self.feed_forward = feed_forward + self.norm1 = paddle.nn.LayerNorm(normalized_shape=size, epsilon=1e-12) + self.norm2 = paddle.nn.LayerNorm(normalized_shape=size, epsilon=1e-12) + self.dropout = paddle.nn.Dropout(p=dropout_rate) + self.size = size + self.normalize_before = normalize_before + + def forward( + self, + x: paddle.Tensor, + mask: paddle.Tensor, + pos_emb: paddle.Tensor, + mask_pad: paddle.Tensor = paddle.ones((0, 0, 0), dtype=paddle.bool), + att_cache: paddle.Tensor = paddle.zeros((0, 0, 0, 0)), + cnn_cache: paddle.Tensor = paddle.zeros((0, 0, 0, 0)), + ) -> Tuple[paddle.Tensor, paddle.Tensor, paddle.Tensor, paddle.Tensor]: + """Compute encoded features. + + Args: + x (torch.Tensor): (#batch, time, size) + mask (torch.Tensor): Mask tensor for the input (#batch, time,time), + (0, 0, 0) means fake mask. + pos_emb (torch.Tensor): just for interface compatibility + to ConformerEncoderLayer + mask_pad (torch.Tensor): does not used in transformer layer, + just for unified api with conformer. + att_cache (torch.Tensor): Cache tensor of the KEY & VALUE + (#batch=1, head, cache_t1, d_k * 2), head * d_k == size. + cnn_cache (torch.Tensor): Convolution cache in conformer layer + (#batch=1, size, cache_t2), not used here, it's for interface + compatibility to ConformerEncoderLayer. + Returns: + torch.Tensor: Output tensor (#batch, time, size). + torch.Tensor: Mask tensor (#batch, time, time). + torch.Tensor: att_cache tensor, + (#batch=1, head, cache_t1 + time, d_k * 2). + torch.Tensor: cnn_cahce tensor (#batch=1, size, cache_t2). + + """ + residual = x + if self.normalize_before: + x = self.norm1(x) + x_att, new_att_cache = self.self_attn( + x, x, x, mask, pos_emb=pos_emb, cache=att_cache + ) + x = residual + self.dropout(x_att) + if not self.normalize_before: + x = self.norm1(x) + residual = x + if self.normalize_before: + x = self.norm2(x) + x = residual + self.dropout(self.feed_forward(x)) + if not self.normalize_before: + x = self.norm2(x) + fake_cnn_cache = paddle.zeros((0, 0, 0), dtype=x.dtype, device=x.place) + return x, mask, new_att_cache, fake_cnn_cache + + +class ConformerEncoderLayer(paddle.nn.Layer): + """Encoder layer module. + Args: + size (int): Input dimension. + self_attn (torch.nn.Module): Self-attention module instance. + `MultiHeadedAttention` or `RelPositionMultiHeadedAttention` + instance can be used as the argument. + feed_forward (torch.nn.Module): Feed-forward module instance. + `PositionwiseFeedForward` instance can be used as the argument. + feed_forward_macaron (torch.nn.Module): Additional feed-forward module + instance. + `PositionwiseFeedForward` instance can be used as the argument. + conv_module (torch.nn.Module): Convolution module instance. + `ConvlutionModule` instance can be used as the argument. + dropout_rate (float): Dropout rate. + normalize_before (bool): + True: use layer_norm before each sub-block. + False: use layer_norm after each sub-block. + """ + + def __init__( + self, + size: int, + self_attn: paddle.nn.Layer, + feed_forward: Optional[paddle.nn.Layer] = None, + feed_forward_macaron: Optional[paddle.nn.Layer] = None, + conv_module: Optional[paddle.nn.Layer] = None, + dropout_rate: float = 0.1, + normalize_before: bool = True, + ): + """Construct an EncoderLayer object.""" + super().__init__() + self.self_attn = self_attn + self.feed_forward = feed_forward + self.feed_forward_macaron = feed_forward_macaron + self.conv_module = conv_module + self.norm_ff = paddle.nn.LayerNorm(normalized_shape=size, epsilon=1e-12) + self.norm_mha = paddle.nn.LayerNorm(normalized_shape=size, epsilon=1e-12) + if feed_forward_macaron is not None: + self.norm_ff_macaron = paddle.nn.LayerNorm( + normalized_shape=size, epsilon=1e-12 + ) + self.ff_scale = 0.5 + else: + self.ff_scale = 1.0 + if self.conv_module is not None: + self.norm_conv = paddle.nn.LayerNorm(normalized_shape=size, epsilon=1e-12) + self.norm_final = paddle.nn.LayerNorm(normalized_shape=size, epsilon=1e-12) + self.dropout = paddle.nn.Dropout(p=dropout_rate) + self.size = size + self.normalize_before = normalize_before + + def forward( + self, + x: paddle.Tensor, + mask: paddle.Tensor, + pos_emb: paddle.Tensor, + mask_pad: paddle.Tensor = paddle.ones((0, 0, 0), dtype=paddle.bool), + att_cache: paddle.Tensor = paddle.zeros((0, 0, 0, 0)), + cnn_cache: paddle.Tensor = paddle.zeros((0, 0, 0, 0)), + ) -> Tuple[paddle.Tensor, paddle.Tensor, paddle.Tensor, paddle.Tensor]: + """Compute encoded features. + + Args: + x (torch.Tensor): (#batch, time, size) + mask (torch.Tensor): Mask tensor for the input (#batch, time,time), + (0, 0, 0) means fake mask. + pos_emb (torch.Tensor): positional encoding, must not be None + for ConformerEncoderLayer. + mask_pad (torch.Tensor): batch padding mask used for conv module. + (#batch, 1,time), (0, 0, 0) means fake mask. + att_cache (torch.Tensor): Cache tensor of the KEY & VALUE + (#batch=1, head, cache_t1, d_k * 2), head * d_k == size. + cnn_cache (torch.Tensor): Convolution cache in conformer layer + (#batch=1, size, cache_t2) + Returns: + torch.Tensor: Output tensor (#batch, time, size). + torch.Tensor: Mask tensor (#batch, time, time). + torch.Tensor: att_cache tensor, + (#batch=1, head, cache_t1 + time, d_k * 2). + torch.Tensor: cnn_cahce tensor (#batch, size, cache_t2). + """ + if self.feed_forward_macaron is not None: + residual = x + if self.normalize_before: + x = self.norm_ff_macaron(x) + x = residual + self.ff_scale * self.dropout(self.feed_forward_macaron(x)) + if not self.normalize_before: + x = self.norm_ff_macaron(x) + residual = x + if self.normalize_before: + x = self.norm_mha(x) + x_att, new_att_cache = self.self_attn(x, x, x, mask, pos_emb, att_cache) + x = residual + self.dropout(x_att) + if not self.normalize_before: + x = self.norm_mha(x) + new_cnn_cache = paddle.zeros((0, 0, 0), dtype=x.dtype, device=x.place) + if self.conv_module is not None: + residual = x + if self.normalize_before: + x = self.norm_conv(x) + x, new_cnn_cache = self.conv_module(x, mask_pad, cnn_cache) + x = residual + self.dropout(x) + if not self.normalize_before: + x = self.norm_conv(x) + residual = x + if self.normalize_before: + x = self.norm_ff(x) + x = residual + self.ff_scale * self.dropout(self.feed_forward(x)) + if not self.normalize_before: + x = self.norm_ff(x) + if self.conv_module is not None: + x = self.norm_final(x) + return x, mask, new_att_cache, new_cnn_cache diff --git a/paddlespeech/t2s/modules/flow/flow.py b/paddlespeech/t2s/modules/flow/flow.py index c22b7a262..2ca92c535 100644 --- a/paddlespeech/t2s/modules/flow/flow.py +++ b/paddlespeech/t2s/modules/flow/flow.py @@ -5,7 +5,7 @@ from typing import Dict, Optional import paddle from omegaconf import DictConfig -from cosyvoice.utils.mask import make_pad_mask +from paddlespeech.t2s.models.CosyVoice.mask import make_pad_mask class MaskedDiffWithXvec(paddle.nn.Layer): @@ -78,7 +78,7 @@ class MaskedDiffWithXvec(paddle.nn.Layer): self.only_mask_loss = only_mask_loss def forward( ->>>>>> self, batch: dict, device: torch.device + self, batch: dict, device: paddle.device ) -> Dict[str, Optional[paddle.Tensor]]: token = batch["speech_token"].to(device) token_len = batch["speech_token_len"].to(device) @@ -88,7 +88,7 @@ class MaskedDiffWithXvec(paddle.nn.Layer): embedding = paddle.nn.functional.normalize(x=embedding, axis=1) embedding = self.spk_embed_affine_layer(embedding) mask = (~make_pad_mask(token_len)).float().unsqueeze(-1).to(device) - token = self.input_embedding(paddle.clamp(token, min=0)) * mask + token = self.input_embedding(paddle.clip(token, min=0)) * mask h, h_lengths = self.encoder(token, token_len) h = self.encoder_proj(h) h, h_lengths = self.length_regulator(h, feat_len) @@ -98,12 +98,12 @@ class MaskedDiffWithXvec(paddle.nn.Layer): continue index = random.randint(0, int(0.3 * j)) conds[i, :index] = feat[i, :index] - conds = conds.transpose(1, 2) + conds = paddle.transpose(conds, perm=[0, 2, 1]) mask = (~make_pad_mask(feat_len)).to(h) loss, _ = self.decoder.compute_loss( - feat.transpose(1, 2).contiguous(), + paddle.transpose(feat, perm=[0, 2, 1]), mask.unsqueeze(1), - h.transpose(1, 2).contiguous(), + paddle.transpose(h, perm=[0, 2, 1]), embedding, cond=conds, ) @@ -130,7 +130,7 @@ class MaskedDiffWithXvec(paddle.nn.Layer): prompt_token_len + token_len, ) mask = (~make_pad_mask(token_len)).unsqueeze(-1).to(embedding) - token = self.input_embedding(paddle.clamp(token, min=0)) * mask + token = self.input_embedding(paddle.clip(token, min=0)) * mask h, h_lengths = self.encoder(token, token_len) h = self.encoder_proj(h) mel_len1, mel_len2 = prompt_feat.shape[1], int( @@ -147,10 +147,11 @@ class MaskedDiffWithXvec(paddle.nn.Layer): [1, mel_len1 + mel_len2, self.output_size], device=token.place ).to(h.dtype) conds[:, :mel_len1] = prompt_feat - conds = conds.transpose(1, 2) + conds = paddle.transpose(conds, perm=[0, 2, 1]) + mask = (~make_pad_mask(paddle.tensor([mel_len1 + mel_len2]))).to(h) feat, flow_cache = self.decoder( - mu=h.transpose(1, 2).contiguous(), + mu=paddle.transpose(h, perm=[0, 2, 1]), mask=mask.unsqueeze(1), spks=embedding, cond=conds, @@ -235,7 +236,7 @@ class CausalMaskedDiffWithXvec(paddle.nn.Layer): self.pre_lookahead_len = pre_lookahead_len def forward( ->>>>>> self, batch: dict, device: torch.device + self, batch: dict, device: paddle.device ) -> Dict[str, Optional[paddle.Tensor]]: token = batch["speech_token"].to(device) token_len = batch["speech_token_len"].to(device) @@ -246,7 +247,7 @@ class CausalMaskedDiffWithXvec(paddle.nn.Layer): embedding = paddle.nn.functional.normalize(x=embedding, axis=1) embedding = self.spk_embed_affine_layer(embedding) mask = (~make_pad_mask(token_len)).float().unsqueeze(-1).to(device) - token = self.input_embedding(paddle.clamp(token, min=0)) * mask + token = self.input_embedding(paddle.clip(token, min=0)) * mask h, h_lengths = self.encoder(token, token_len, streaming=streaming) h = self.encoder_proj(h) conds = paddle.zeros(feat.shape, device=token.place) @@ -255,12 +256,13 @@ class CausalMaskedDiffWithXvec(paddle.nn.Layer): continue index = random.randint(0, int(0.3 * j)) conds[i, :index] = feat[i, :index] - conds = conds.transpose(1, 2) + conds = paddle.transpose(conds, perm=[0, 2, 1]) + mask = (~make_pad_mask(h_lengths.sum(dim=-1).squeeze(dim=1))).to(h) loss, _ = self.decoder.compute_loss( - feat.transpose(1, 2).contiguous(), + paddle.transpose(feat, perm=[0, 2, 1]).contiguous(), mask.unsqueeze(1), - h.transpose(1, 2).contiguous(), + paddle.transpose(h, perm=[0, 2, 1]).contiguous(), embedding, cond=conds, streaming=streaming, @@ -283,12 +285,13 @@ class CausalMaskedDiffWithXvec(paddle.nn.Layer): assert token.shape[0] == 1 embedding = paddle.nn.functional.normalize(x=embedding, axis=1) embedding = self.spk_embed_affine_layer(embedding) + token, token_len = ( paddle.cat([prompt_token, token], dim=1), prompt_token_len + token_len, ) mask = (~make_pad_mask(token_len)).unsqueeze(-1).to(embedding) - token = self.input_embedding(paddle.clamp(token, min=0)) * mask + token = self.input_embedding(paddle.clip(token, min=0)) * mask if finalize is True: h, h_lengths = self.encoder(token, token_len, streaming=streaming) else: @@ -302,19 +305,20 @@ class CausalMaskedDiffWithXvec(paddle.nn.Layer): mel_len1, mel_len2 = prompt_feat.shape[1], h.shape[1] - prompt_feat.shape[1] h = self.encoder_proj(h) conds = paddle.zeros( - [1, mel_len1 + mel_len2, self.output_size], device=token.place + [1, mel_len1 + mel_len2, self.output_size] ).to(h.dtype) conds[:, :mel_len1] = prompt_feat - conds = conds.transpose(1, 2) - mask = (~make_pad_mask(paddle.tensor([mel_len1 + mel_len2]))).to(h) + conds = paddle.transpose(conds, perm=[0, 2, 1]) + mask = (~make_pad_mask(paddle.to_tensor([mel_len1 + mel_len2],dtype='int32'))).to(h) feat, _ = self.decoder( - mu=h.transpose(1, 2).contiguous(), + mu=paddle.transpose(h, perm=[0, 2, 1]).contiguous(), mask=mask.unsqueeze(1), spks=embedding, cond=conds, n_timesteps=10, streaming=streaming, ) + paddle.save(feat,'/root/paddlejob/workspace/zhangjinghong/CosyVoice/feat.pdparams') feat = feat[:, :, mel_len1:] assert feat.shape[2] == mel_len2 return feat.float(), None diff --git a/paddlespeech/t2s/modules/flow/flow_matching.py b/paddlespeech/t2s/modules/flow/flow_matching.py index bf8c9f7a0..9c68deb48 100644 --- a/paddlespeech/t2s/modules/flow/flow_matching.py +++ b/paddlespeech/t2s/modules/flow/flow_matching.py @@ -1,8 +1,100 @@ import paddle -from matcha.models.components.flow_matching import BASECFM +from abc import ABC +from paddlespeech.t2s.models.CosyVoice.common import set_all_random_seed -from cosyvoice.utils.common import set_all_random_seed +class BASECFM(paddle.nn.Layer, ABC): + def __init__(self, n_feats, cfm_params, n_spks=1, spk_emb_dim=128): + super().__init__() + self.n_feats = n_feats + self.n_spks = n_spks + self.spk_emb_dim = spk_emb_dim + self.solver = cfm_params.solver + if hasattr(cfm_params, "sigma_min"): + self.sigma_min = cfm_params.sigma_min + else: + self.sigma_min = 0.0001 + self.estimator = None + + @paddle.no_grad() + def forward(self, mu, mask, n_timesteps, temperature=1.0, spks=None, cond=None): + """Forward diffusion + + Args: + mu (torch.Tensor): output of encoder + shape: (batch_size, n_feats, mel_timesteps) + mask (torch.Tensor): output_mask + shape: (batch_size, 1, mel_timesteps) + n_timesteps (int): number of diffusion steps + temperature (float, optional): temperature for scaling noise. Defaults to 1.0. + spks (torch.Tensor, optional): speaker ids. Defaults to None. + shape: (batch_size, spk_emb_dim) + cond: Not used but kept for future purposes + Returns: + sample: generated mel-spectrogram + shape: (batch_size, n_feats, mel_timesteps) + """ + z = paddle.randn(shape=mu.shape, dtype=mu.dtype) * temperature + t_span = paddle.linspace(start=0, stop=1, num=n_timesteps + 1) + return self.solve_euler( + z, t_span=t_span, mu=mu, mask=mask, spks=spks, cond=cond + ) + + def solve_euler(self, x, t_span, mu, mask, spks, cond): + """ + Fixed euler solver for ODEs. + Args: + x (torch.Tensor): random noise + t_span (torch.Tensor): n_timesteps interpolated + shape: (n_timesteps + 1,) + mu (torch.Tensor): output of encoder + shape: (batch_size, n_feats, mel_timesteps) + mask (torch.Tensor): output_mask + shape: (batch_size, 1, mel_timesteps) + spks (torch.Tensor, optional): speaker ids. Defaults to None. + shape: (batch_size, spk_emb_dim) + cond: Not used but kept for future purposes + """ + t, _, dt = t_span[0], t_span[-1], t_span[1] - t_span[0] + sol = [] + for step in range(1, len(t_span)): + dphi_dt = self.estimator(x, mask, mu, t, spks, cond) + x = x + dt * dphi_dt + t = t + dt + sol.append(x) + if step < len(t_span) - 1: + dt = t_span[step + 1] - t + return sol[-1] + + def compute_loss(self, x1, mask, mu, spks=None, cond=None): + """Computes diffusion loss + + Args: + x1 (torch.Tensor): Target + shape: (batch_size, n_feats, mel_timesteps) + mask (torch.Tensor): target mask + shape: (batch_size, 1, mel_timesteps) + mu (torch.Tensor): output of encoder + shape: (batch_size, n_feats, mel_timesteps) + spks (torch.Tensor, optional): speaker embedding. Defaults to None. + shape: (batch_size, spk_emb_dim) + + Returns: + loss: conditional flow matching loss + y: conditional flow + shape: (batch_size, n_feats, mel_timesteps) + """ + b, _, t = mu.shape + t = paddle.rand(shape=[b, 1, 1], dtype=mu.dtype) + z = paddle.randn(shape=x1.shape, dtype=x1.dtype) + y = (1 - (1 - self.sigma_min) * t) * z + t * x1 + u = x1 - (1 - self.sigma_min) * z + loss = paddle.nn.functional.mse_loss( + input=self.estimator(y, mask, mu, t.squeeze(), spks), + label=u, + reduction="sum", + ) / (paddle.sum(mask) * u.shape[1]) + return loss, y class ConditionalCFM(BASECFM): def __init__( @@ -35,7 +127,7 @@ class ConditionalCFM(BASECFM): spks=None, cond=None, prompt_len=0, - cache=paddle.zeros(1, 80, 0, 2), + cache=paddle.zeros([1, 80, 0, 2]), ): """Forward diffusion @@ -62,9 +154,9 @@ class ConditionalCFM(BASECFM): if cache_size != 0: z[:, :, :cache_size] = cache[:, :, :, 0] mu[:, :, :cache_size] = cache[:, :, :, 1] - z_cache = paddle.cat([z[:, :, :prompt_len], z[:, :, -34:]], dim=2) - mu_cache = paddle.cat([mu[:, :, :prompt_len], mu[:, :, -34:]], dim=2) - cache = paddle.stack([z_cache, mu_cache], dim=-1) + z_cache = paddle.cat([z[:, :, :prompt_len], z[:, :, -34:]], axis=2) + mu_cache = paddle.cat([mu[:, :, :prompt_len], mu[:, :, -34:]], axis=2) + cache = paddle.stack([z_cache, mu_cache], axis=-1) t_span = paddle.linspace(start=0, stop=1, num=n_timesteps + 1, dtype=mu.dtype) if self.t_scheduler == "cosine": t_span = 1 - paddle.cos(t_span * 0.5 * paddle.pi) @@ -89,14 +181,14 @@ class ConditionalCFM(BASECFM): cond: Not used but kept for future purposes """ t, _, dt = t_span[0], t_span[-1], t_span[1] - t_span[0] - t = t.unsqueeze(dim=0) + t = t.unsqueeze(axis=0) sol = [] - x_in = paddle.zeros([2, 80, x.size(2)], device=x.place, dtype=x.dtype) - mask_in = paddle.zeros([2, 1, x.size(2)], device=x.place, dtype=x.dtype) - mu_in = paddle.zeros([2, 80, x.size(2)], device=x.place, dtype=x.dtype) - t_in = paddle.zeros([2], device=x.place, dtype=x.dtype) - spks_in = paddle.zeros([2, 80], device=x.place, dtype=x.dtype) - cond_in = paddle.zeros([2, 80, x.size(2)], device=x.place, dtype=x.dtype) + x_in = paddle.zeros([2, 80, x.shape[2]], dtype=x.dtype) + mask_in = paddle.zeros([2, 1, x.shape[2]], dtype=x.dtype) + mu_in = paddle.zeros([2, 80, x.shape[2]], dtype=x.dtype) + t_in = paddle.zeros([2], dtype=x.dtype) + spks_in = paddle.zeros([2, 80], dtype=x.dtype) + cond_in = paddle.zeros([2, 80, x.shape[2]], dtype=x.dtype) for step in range(1, len(t_span)): x_in[:] = x mask_in[:] = mask @@ -107,9 +199,7 @@ class ConditionalCFM(BASECFM): dphi_dt = self.forward_estimator( x_in, mask_in, mu_in, t_in, spks_in, cond_in, streaming ) - dphi_dt, cfg_dphi_dt = paddle.compat.split( - dphi_dt, [x.size(0), x.size(0)], dim=0 - ) + dphi_dt, cfg_dphi_dt = paddle.split(dphi_dt, [x.shape[0], x.shape[0]], axis=0) dphi_dt = ( 1.0 + self.inference_cfg_rate ) * dphi_dt - self.inference_cfg_rate * cfg_dphi_dt @@ -127,12 +217,12 @@ class ConditionalCFM(BASECFM): [estimator, stream], trt_engine = self.estimator.acquire_estimator() paddle.device.current_stream().synchronize() with stream: - estimator.set_input_shape("x", (2, 80, x.size(2))) - estimator.set_input_shape("mask", (2, 1, x.size(2))) - estimator.set_input_shape("mu", (2, 80, x.size(2))) + estimator.set_input_shape("x", (2, 80, x.shape[2])) + estimator.set_input_shape("mask", (2, 1, x.shape[2])) + estimator.set_input_shape("mu", (2, 80, x.shape[2])) estimator.set_input_shape("t", (2,)) estimator.set_input_shape("spks", (2, 80)) - estimator.set_input_shape("cond", (2, 80, x.size(2))) + estimator.set_input_shape("cond", (2, 80, x.shape[2])) data_ptrs = [ x.contiguous().data_ptr(), mask.contiguous().data_ptr(), @@ -201,9 +291,8 @@ class CausalConditionalCFM(ConditionalCFM): estimator: paddle.nn.Layer = None, ): super().__init__(in_channels, cfm_params, n_spks, spk_emb_dim, estimator) - set_all_random_seed(0) + set_all_random_seed(42) self.rand_noise = paddle.randn([1, 80, 50 * 300]) - @paddle.no_grad() def forward( self, @@ -232,7 +321,9 @@ class CausalConditionalCFM(ConditionalCFM): sample: generated mel-spectrogram shape: (batch_size, n_feats, mel_timesteps) """ - z = self.rand_noise[:, :, : mu.size(2)].to(mu.place).to(mu.dtype) * temperature + + z = self.rand_noise[:, :, : mu.shape[2]].to(mu.place).to(mu.dtype) * temperature + t_span = paddle.linspace(start=0, stop=1, num=n_timesteps + 1, dtype=mu.dtype) if self.t_scheduler == "cosine": t_span = 1 - paddle.cos(t_span * 0.5 * paddle.pi) diff --git a/paddlespeech/t2s/modules/flow/flow_matching_matcha.py b/paddlespeech/t2s/modules/flow/flow_matching_matcha.py new file mode 100644 index 000000000..0bcea69e4 --- /dev/null +++ b/paddlespeech/t2s/modules/flow/flow_matching_matcha.py @@ -0,0 +1,124 @@ +from abc import ABC + +import paddle +from matcha.models.components.decoder import Decoder +from matcha.utils.pylogger import get_pylogger + +log = get_pylogger(__name__) + + +class BASECFM(paddle.nn.Layer, ABC): + def __init__(self, n_feats, cfm_params, n_spks=1, spk_emb_dim=128): + super().__init__() + self.n_feats = n_feats + self.n_spks = n_spks + self.spk_emb_dim = spk_emb_dim + self.solver = cfm_params.solver + if hasattr(cfm_params, "sigma_min"): + self.sigma_min = cfm_params.sigma_min + else: + self.sigma_min = 0.0001 + self.estimator = None + + @paddle.no_grad() + def forward(self, mu, mask, n_timesteps, temperature=1.0, spks=None, cond=None): + """Forward diffusion + + Args: + mu (torch.Tensor): output of encoder + shape: (batch_size, n_feats, mel_timesteps) + mask (torch.Tensor): output_mask + shape: (batch_size, 1, mel_timesteps) + n_timesteps (int): number of diffusion steps + temperature (float, optional): temperature for scaling noise. Defaults to 1.0. + spks (torch.Tensor, optional): speaker ids. Defaults to None. + shape: (batch_size, spk_emb_dim) + cond: Not used but kept for future purposes + + Returns: + sample: generated mel-spectrogram + shape: (batch_size, n_feats, mel_timesteps) + """ + z = paddle.randn(shape=mu.shape, dtype=mu.dtype) * temperature + t_span = paddle.linspace(start=0, stop=1, num=n_timesteps + 1) + return self.solve_euler( + z, t_span=t_span, mu=mu, mask=mask, spks=spks, cond=cond + ) + + def solve_euler(self, x, t_span, mu, mask, spks, cond): + """ + Fixed euler solver for ODEs. + Args: + x (torch.Tensor): random noise + t_span (torch.Tensor): n_timesteps interpolated + shape: (n_timesteps + 1,) + mu (torch.Tensor): output of encoder + shape: (batch_size, n_feats, mel_timesteps) + mask (torch.Tensor): output_mask + shape: (batch_size, 1, mel_timesteps) + spks (torch.Tensor, optional): speaker ids. Defaults to None. + shape: (batch_size, spk_emb_dim) + cond: Not used but kept for future purposes + """ + t, _, dt = t_span[0], t_span[-1], t_span[1] - t_span[0] + sol = [] + for step in range(1, len(t_span)): + dphi_dt = self.estimator(x, mask, mu, t, spks, cond) + x = x + dt * dphi_dt + t = t + dt + sol.append(x) + if step < len(t_span) - 1: + dt = t_span[step + 1] - t + return sol[-1] + + def compute_loss(self, x1, mask, mu, spks=None, cond=None): + """Computes diffusion loss + + Args: + x1 (torch.Tensor): Target + shape: (batch_size, n_feats, mel_timesteps) + mask (torch.Tensor): target mask + shape: (batch_size, 1, mel_timesteps) + mu (torch.Tensor): output of encoder + shape: (batch_size, n_feats, mel_timesteps) + spks (torch.Tensor, optional): speaker embedding. Defaults to None. + shape: (batch_size, spk_emb_dim) + + Returns: + loss: conditional flow matching loss + y: conditional flow + shape: (batch_size, n_feats, mel_timesteps) + """ + b, _, t = mu.shape + t = paddle.rand(shape=[b, 1, 1], dtype=mu.dtype) + z = paddle.randn(shape=x1.shape, dtype=x1.dtype) + y = (1 - (1 - self.sigma_min) * t) * z + t * x1 + u = x1 - (1 - self.sigma_min) * z + loss = paddle.nn.functional.mse_loss( + input=self.estimator(y, mask, mu, t.squeeze(), spks), + label=u, + reduction="sum", + ) / (paddle.sum(mask) * u.shape[1]) + return loss, y + + +class CFM(BASECFM): + def __init__( + self, + in_channels, + out_channel, + cfm_params, + decoder_params, + n_spks=1, + spk_emb_dim=64, + ): + super().__init__( + n_feats=in_channels, + cfm_params=cfm_params, + n_spks=n_spks, + spk_emb_dim=spk_emb_dim, + ) + in_channels = in_channels + (spk_emb_dim if n_spks > 1 else 0) + self.estimator = Decoder( + in_channels=in_channels, out_channels=out_channel, **decoder_params + ) diff --git a/paddlespeech/t2s/modules/flow/lora.py b/paddlespeech/t2s/modules/flow/lora.py new file mode 100644 index 000000000..557f9ff15 --- /dev/null +++ b/paddlespeech/t2s/modules/flow/lora.py @@ -0,0 +1,123 @@ +from typing import Optional, Tuple, Union + +import paddle + + +class LoRALinearLayer(paddle.nn.Layer): + """ + A linear layer that is used with LoRA. + + Parameters: + in_features (`int`): + Number of input features. + out_features (`int`): + Number of output features. + rank (`int`, `optional`, defaults to 4): + The rank of the LoRA layer. + network_alpha (`float`, `optional`, defaults to `None`): + The value of the network alpha used for stable learning and preventing underflow. This value has the same + meaning as the `--network_alpha` option in the kohya-ss trainer script. See + https://github.com/darkstorm2150/sd-scripts/blob/main/docs/train_network_README-en.md#execute-learning + device (`torch.device`, `optional`, defaults to `None`): + The device to use for the layer's weights. + dtype (`torch.dtype`, `optional`, defaults to `None`): + The dtype to use for the layer's weights. + """ + + def __init__( + self, + in_features: int, + out_features: int, + rank: int = 4, + network_alpha: Optional[float] = None, + device: Optional[Union[paddle.CPUPlace, paddle.CUDAPlace, str]] = None, + dtype: Optional[paddle.dtype] = None, + ): + super().__init__() + self.down = paddle.nn.Linear( + in_features=in_features, out_features=rank, bias_attr=False + ) + self.up = paddle.nn.Linear( + in_features=rank, out_features=out_features, bias_attr=False + ) + self.network_alpha = network_alpha + self.rank = rank + self.out_features = out_features + self.in_features = in_features + paddle.nn.init.normal_(self.down.weight, std=1 / rank) + paddle.nn.init.zeros_(self.up.weight) + + def forward(self, hidden_states: paddle.Tensor) -> paddle.Tensor: + orig_dtype = hidden_states.dtype + dtype = self.down.weight.dtype + down_hidden_states = self.down(hidden_states.to(dtype)) + up_hidden_states = self.up(down_hidden_states) + if self.network_alpha is not None: + up_hidden_states *= self.network_alpha / self.rank + return up_hidden_states.to(orig_dtype) + + + +class LoRACompatibleLinear(paddle.nn.Linear): + """ + A Linear layer that can be used with LoRA. + """ + + def __init__(self, *args, lora_layer: Optional[LoRALinearLayer] = None, **kwargs): + super().__init__(*args, **kwargs) + self.lora_layer = lora_layer + + def set_lora_layer(self, lora_layer: Optional[LoRALinearLayer]): + self.lora_layer = lora_layer + + def _fuse_lora(self, lora_scale: float = 1.0, safe_fusing: bool = False): + if self.lora_layer is None: + return + dtype, device = self.weight.data.dtype, self.weight.data.place + w_orig = self.weight.data.float() + w_up = self.lora_layer.up.weight.data.float() + w_down = self.lora_layer.down.weight.data.float() + if self.lora_layer.network_alpha is not None: + w_up = w_up * self.lora_layer.network_alpha / self.lora_layer.rank + fused_weight = ( + w_orig + lora_scale * paddle.bmm(w_up[None, :], w_down[None, :])[0] + ) + if safe_fusing and paddle.isnan(fused_weight).any().item(): + raise ValueError( + f"This LoRA weight seems to be broken. Encountered NaN values when trying to fuse LoRA weights for {self}.LoRA weights will not be fused." + ) + self.weight.data = fused_weight.to(device=device, dtype=dtype) + self.lora_layer = None + self.w_up = w_up.cpu() + self.w_down = w_down.cpu() + self._lora_scale = lora_scale + + def _unfuse_lora(self): + if not ( + getattr(self, "w_up", None) is not None + and getattr(self, "w_down", None) is not None + ): + return + fused_weight = self.weight.data + dtype, device = fused_weight.dtype, fused_weight.place + w_up = self.w_up.to(device=device).float() + w_down = self.w_down.to(device).float() + unfused_weight = ( + fused_weight.float() + - self._lora_scale * paddle.bmm(w_up[None, :], w_down[None, :])[0] + ) + self.weight.data = unfused_weight.to(device=device, dtype=dtype) + self.w_up = None + self.w_down = None + + def forward( + self, hidden_states: paddle.Tensor, scale: float = 1.0 + ) -> paddle.Tensor: + if self.lora_layer is None: + out = super().forward(hidden_states) + return out + else: + out = super().forward(hidden_states) + scale * self.lora_layer( + hidden_states + ) + return out diff --git a/paddlespeech/t2s/modules/flow/mask.py b/paddlespeech/t2s/modules/flow/mask.py new file mode 100644 index 000000000..697ca0eb5 --- /dev/null +++ b/paddlespeech/t2s/modules/flow/mask.py @@ -0,0 +1,287 @@ +import paddle + +def device2str(type=None, index=None, *, device=None): + type = device if device else type + if isinstance(type, int): + type = f'gpu:{type}' + elif isinstance(type, str): + if 'cuda' in type: + type = type.replace('cuda', 'gpu') + if 'cpu' in type: + type = 'cpu' + elif index is not None: + type = f'{type}:{index}' + elif isinstance(type, paddle.CPUPlace) or (type is None): + type = 'cpu' + elif isinstance(type, paddle.CUDAPlace): + type = f'gpu:{type.get_device_id()}' + + return type + +def _Tensor_max(self, *args, **kwargs): + if "other" in kwargs: + kwargs["y"] = kwargs.pop("other") + ret = paddle.maximum(self, *args, **kwargs) + elif len(args) == 1 and isinstance(args[0], paddle.Tensor): + ret = paddle.maximum(self, *args, **kwargs) + else: + if "dim" in kwargs: + kwargs["axis"] = kwargs.pop("dim") + + if "axis" in kwargs or len(args) >= 1: + ret = paddle.max(self, *args, **kwargs), paddle.argmax(self, *args, **kwargs) + else: + ret = paddle.max(self, *args, **kwargs) + + return ret + +setattr(paddle.Tensor, "_max", _Tensor_max) + + + +""" +def subsequent_mask( + size: int, + device: torch.device = torch.device("cpu"), +) -> torch.Tensor: + ""\"Create mask for subsequent steps (size, size). + + This mask is used only in decoder which works in an auto-regressive mode. + This means the current step could only do attention with its left steps. + + In encoder, fully attention is used when streaming is not necessary and + the sequence is not long. In this case, no attention mask is needed. + + When streaming is need, chunk-based attention is used in encoder. See + subsequent_chunk_mask for the chunk-based attention mask. + + Args: + size (int): size of mask + str device (str): "cpu" or "cuda" or torch.Tensor.device + dtype (torch.device): result dtype + + Returns: + torch.Tensor: mask + + Examples: + >>> subsequent_mask(3) + [[1, 0, 0], + [1, 1, 0], + [1, 1, 1]] + ""\" + ret = torch.ones(size, size, device=device, dtype=torch.bool) + return torch.tril(ret) +""" + + +def subsequent_mask( +>>>>>> size: int, device: torch.device = device2str("cpu") +) -> paddle.Tensor: + """Create mask for subsequent steps (size, size). + + This mask is used only in decoder which works in an auto-regressive mode. + This means the current step could only do attention with its left steps. + + In encoder, fully attention is used when streaming isnot necessary and + the sequence is not long. In this case, no attention mask is needed. + + When streaming is need, chunk-based attention is used in encoder. See + subsequent_chunk_mask for the chunk-based attention mask. + + Args: + size (int): size of mask + str device (str): "cpu" or "cuda" or torch.Tensor.device + dtype (torch.device): result dtype + + Returns: + torch.Tensor: mask + + Examples: + >>> subsequent_mask(3) + [[1, 0, 0], + [1, 1, 0], + [1, 1, 1]] + """ + arange = paddle.arange(size, device=device) + mask = arange.expand(size, size) + arange = arange.unsqueeze(-1) + mask = mask <= arange + return mask + + +def subsequent_chunk_mask_deprecated( + size: int, + chunk_size: int, + num_left_chunks: int = -1, +>>>>>> device: torch.device = device2str("cpu"), +) -> paddle.Tensor: + """Create mask for subsequent steps (size, size) with chunk size, + this is for streaming encoder + + Args: + size (int): size of mask + chunk_size (int): size of chunk + num_left_chunks (int): number of left chunks + <0: use full chunk + >=0: use num_left_chunks + device (torch.device): "cpu" or "cuda" or torch.Tensor.device + + Returns: + torch.Tensor: mask + + Examples: + >>> subsequent_chunk_mask(4, 2) + [[1, 1, 0, 0], + [1, 1, 0, 0], + [1, 1, 1, 1], + [1, 1, 1, 1]] + """ + ret = paddle.zeros(size, size, device=device, dtype=paddle.bool) + for i in range(size): + if num_left_chunks < 0: + start = 0 + else: + start = max((i // chunk_size - num_left_chunks) * chunk_size, 0) + ending = min((i // chunk_size + 1) * chunk_size, size) + ret[i, start:ending] = True + return ret + + +def subsequent_chunk_mask( + size: int, + chunk_size: int, + num_left_chunks: int = -1, +>>>>>> device: torch.device = device2str("cpu"), +) -> paddle.Tensor: + """Create mask for subsequent steps (size, size) with chunk size, + this is for streaming encoder + + Args: + size (int): size of mask + chunk_size (int): size of chunk + num_left_chunks (int): number of left chunks + <0: use full chunk + >=0: use num_left_chunks + device (torch.device): "cpu" or "cuda" or torch.Tensor.device + + Returns: + torch.Tensor: mask + + Examples: + >>> subsequent_chunk_mask(4, 2) + [[1, 1, 0, 0], + [1, 1, 0, 0], + [1, 1, 1, 1], + [1, 1, 1, 1]] + """ + pos_idx = paddle.arange(size, device=device) + block_value = ( + paddle.div(pos_idx, chunk_size, rounding_mode="trunc") + 1 + ) * chunk_size + ret = pos_idx.unsqueeze(0) < block_value.unsqueeze(1) + return ret + + +def add_optional_chunk_mask( + xs: paddle.Tensor, + masks: paddle.Tensor, + use_dynamic_chunk: bool, + use_dynamic_left_chunk: bool, + decoding_chunk_size: int, + static_chunk_size: int, + num_decoding_left_chunks: int, + enable_full_context: bool = True, +): + """Apply optional mask for encoder. + + Args: + xs (torch.Tensor): padded input, (B, L, D), L for max length + mask (torch.Tensor): mask for xs, (B, 1, L) + use_dynamic_chunk (bool): whether to use dynamic chunk or not + use_dynamic_left_chunk (bool): whether to use dynamic left chunk for + training. + decoding_chunk_size (int): decoding chunk size for dynamic chunk, it's + 0: default for training, use random dynamic chunk. + <0: for decoding, use full chunk. + >0: for decoding, use fixed chunk size as set. + static_chunk_size (int): chunk size for static chunk training/decoding + if it's greater than 0, if use_dynamic_chunk is true, + this parameter will be ignored + num_decoding_left_chunks: number of left chunks, this is for decoding, + the chunk size is decoding_chunk_size. + >=0: use num_decoding_left_chunks + <0: use all left chunks + enable_full_context (bool): + True: chunk size is either [1, 25] or full context(max_len) + False: chunk size ~ U[1, 25] + + Returns: + torch.Tensor: chunk mask of the input xs. + """ + if use_dynamic_chunk: + max_len = xs.size(1) + if decoding_chunk_size < 0: + chunk_size = max_len + num_left_chunks = -1 + elif decoding_chunk_size > 0: + chunk_size = decoding_chunk_size + num_left_chunks = num_decoding_left_chunks + else: + chunk_size = paddle.randint(low=1, high=max_len, shape=(1,)).item() + num_left_chunks = -1 + if chunk_size > max_len // 2 and enable_full_context: + chunk_size = max_len + else: + chunk_size = chunk_size % 25 + 1 + if use_dynamic_left_chunk: + max_left_chunks = (max_len - 1) // chunk_size + num_left_chunks = paddle.randint( + low=0, high=max_left_chunks, shape=(1,) + ).item() + chunk_masks = subsequent_chunk_mask( + xs.size(1), chunk_size, num_left_chunks, xs.place + ) + chunk_masks = chunk_masks.unsqueeze(0) + chunk_masks = masks & chunk_masks + elif static_chunk_size > 0: + num_left_chunks = num_decoding_left_chunks + chunk_masks = subsequent_chunk_mask( + xs.size(1), static_chunk_size, num_left_chunks, xs.place + ) + chunk_masks = chunk_masks.unsqueeze(0) + chunk_masks = masks & chunk_masks + else: + chunk_masks = masks + assert chunk_masks.dtype == paddle.bool + if (chunk_masks.sum(dim=-1) == 0).sum().item() != 0: + print( + "get chunk_masks all false at some timestep, force set to true, make sure they are masked in futuer computation!" + ) + chunk_masks[chunk_masks.sum(dim=-1) == 0] = True + return chunk_masks + + +def make_pad_mask(lengths: paddle.Tensor, max_len: int = 0) -> paddle.Tensor: + """Make mask tensor containing indices of padded part. + + See description of make_non_pad_mask. + + Args: + lengths (torch.Tensor): Batch of lengths (B,). + Returns: + torch.Tensor: Mask tensor containing indices of padded part. + + Examples: + >>> lengths = [5, 3, 2] + >>> make_pad_mask(lengths) + masks = [[0, 0, 0, 0 ,0], + [0, 0, 0, 1, 1], + [0, 0, 1, 1, 1]] + """ + batch_size = lengths.size(0) + max_len = max_len if max_len > 0 else lengths._max().item() + seq_range = paddle.arange(0, max_len, dtype=paddle.int64, device=lengths.place) + seq_range_expand = seq_range.unsqueeze(0).expand(batch_size, max_len) + seq_length_expand = lengths.unsqueeze(-1) + mask = seq_range_expand >= seq_length_expand + return mask \ No newline at end of file diff --git a/paddlespeech/t2s/modules/flow/matcha_decoder.py b/paddlespeech/t2s/modules/flow/matcha_decoder.py new file mode 100644 index 000000000..7ecb100b2 --- /dev/null +++ b/paddlespeech/t2s/modules/flow/matcha_decoder.py @@ -0,0 +1,445 @@ +import math +from typing import Optional + +import einops +import paddle +from conformer import ConformerBlock +from matcha.models.components.transformer import BasicTransformerBlock +ACTIVATION_FUNCTIONS = { + "swish": paddle.nn.SiLU(), + "silu": paddle.nn.SiLU(), + "mish": paddle.nn.Mish(), + "gelu": paddle.nn.GELU(), + "relu": paddle.nn.ReLU(), +} +def get_activation(act_fn: str) -> paddle.nn.Layer: + """Helper function to get activation function from string. + + Args: + act_fn (str): Name of activation function. + + Returns: + nn.Module: Activation function. + """ + act_fn = act_fn.lower() + if act_fn in ACTIVATION_FUNCTIONS: + return ACTIVATION_FUNCTIONS[act_fn] + else: + raise ValueError(f"Unsupported activation function: {act_fn}") + +class SinusoidalPosEmb(paddle.nn.Layer): + def __init__(self, dim): + super().__init__() + self.dim = dim + assert self.dim % 2 == 0, "SinusoidalPosEmb requires dim to be even" + + def forward(self, x, scale=1000): + if x.ndim < 1: + x = x.unsqueeze(0) + device = x.place + half_dim = self.dim // 2 + emb = math.log(10000) / (half_dim - 1) + emb = paddle.exp(x=paddle.arange(half_dim, device=device).float() * -emb) + emb = scale * x.unsqueeze(1) * emb.unsqueeze(0) + emb = paddle.cat((emb.sin(), emb.cos()), dim=-1) + return emb + + +class Block1D(paddle.nn.Layer): + def __init__(self, dim, dim_out, groups=8): + super().__init__() + self.block = paddle.nn.Sequential( + paddle.nn.Conv1d(dim, dim_out, 3, padding=1), + paddle.nn.GroupNorm(num_groups=groups, num_channels=dim_out), + paddle.nn.Mish(), + ) + + def forward(self, x, mask): + output = self.block(x * mask) + return output * mask + + +class ResnetBlock1D(paddle.nn.Layer): + def __init__(self, dim, dim_out, time_emb_dim, groups=8): + super().__init__() + self.mlp = paddle.nn.Sequential( + paddle.nn.Mish(), + paddle.nn.Linear(in_features=time_emb_dim, out_features=dim_out), + ) + self.block1 = Block1D(dim, dim_out, groups=groups) + self.block2 = Block1D(dim_out, dim_out, groups=groups) + self.res_conv = paddle.nn.Conv1d(dim, dim_out, 1) + + def forward(self, x, mask, time_emb): + h = self.block1(x, mask) + h += self.mlp(time_emb).unsqueeze(-1) + h = self.block2(h, mask) + output = h + self.res_conv(x * mask) + return output + + +class Downsample1D(paddle.nn.Layer): + def __init__(self, dim): + super().__init__() + self.conv = paddle.nn.Conv1d(dim, dim, 3, 2, 1) + + def forward(self, x): + return self.conv(x) + + +class TimestepEmbedding(paddle.nn.Layer): + def __init__( + self, + in_channels: int, + time_embed_dim: int, + act_fn: str = "silu", + out_dim: int = None, + post_act_fn: Optional[str] = None, + cond_proj_dim=None, + ): + super().__init__() + self.linear_1 = paddle.nn.Linear( + in_features=in_channels, out_features=time_embed_dim + ) + if cond_proj_dim is not None: + self.cond_proj = paddle.nn.Linear( + in_features=cond_proj_dim, out_features=in_channels, bias_attr=False + ) + else: + self.cond_proj = None + self.act = get_activation(act_fn) + if out_dim is not None: + time_embed_dim_out = out_dim + else: + time_embed_dim_out = time_embed_dim + self.linear_2 = paddle.nn.Linear( + in_features=time_embed_dim, out_features=time_embed_dim_out + ) + if post_act_fn is None: + self.post_act = None + else: + self.post_act = get_activation(post_act_fn) + + def forward(self, sample, condition=None): + if condition is not None: + sample = sample + self.cond_proj(condition) + sample = self.linear_1(sample) + if self.act is not None: + sample = self.act(sample) + sample = self.linear_2(sample) + if self.post_act is not None: + sample = self.post_act(sample) + return sample + + +class Upsample1D(paddle.nn.Layer): + """A 1D upsampling layer with an optional convolution. + + Parameters: + channels (`int`): + number of channels in the inputs and outputs. + use_conv (`bool`, default `False`): + option to use a convolution. + use_conv_transpose (`bool`, default `False`): + option to use a convolution transpose. + out_channels (`int`, optional): + number of output channels. Defaults to `channels`. + """ + + def __init__( + self, + channels, + use_conv=False, + use_conv_transpose=True, + out_channels=None, + name="conv", + ): + super().__init__() + self.channels = channels + self.out_channels = out_channels or channels + self.use_conv = use_conv + self.use_conv_transpose = use_conv_transpose + self.name = name + self.conv = None + if use_conv_transpose: + self.conv = paddle.nn.Conv1DTranspose( + in_channels=channels, + out_channels=self.out_channels, + kernel_size=4, + stride=2, + padding=1, + ) + elif use_conv: + self.conv = paddle.nn.Conv1d(self.channels, self.out_channels, 3, padding=1) + + def forward(self, inputs): + assert inputs.shape[1] == self.channels + if self.use_conv_transpose: + return self.conv(inputs) + outputs = paddle.nn.functional.interpolate( + x=inputs, scale_factor=2.0, mode="nearest" + ) + if self.use_conv: + outputs = self.conv(outputs) + return outputs + + +class ConformerWrapper(ConformerBlock): + def __init__( + self, + *, + dim, + dim_head=64, + heads=8, + ff_mult=4, + conv_expansion_factor=2, + conv_kernel_size=31, + attn_dropout=0, + ff_dropout=0, + conv_dropout=0, + conv_causal=False, + ): + super().__init__( + dim=dim, + dim_head=dim_head, + heads=heads, + ff_mult=ff_mult, + conv_expansion_factor=conv_expansion_factor, + conv_kernel_size=conv_kernel_size, + attn_dropout=attn_dropout, + ff_dropout=ff_dropout, + conv_dropout=conv_dropout, + conv_causal=conv_causal, + ) + + def forward( + self, + hidden_states, + attention_mask, + encoder_hidden_states=None, + encoder_attention_mask=None, + timestep=None, + ): + return super().forward(x=hidden_states, mask=attention_mask.bool()) + + +class Decoder(paddle.nn.Layer): + def __init__( + self, + in_channels, + out_channels, + channels=(256, 256), + dropout=0.05, + attention_head_dim=64, + n_blocks=1, + num_mid_blocks=2, + num_heads=4, + act_fn="snake", + down_block_type="transformer", + mid_block_type="transformer", + up_block_type="transformer", + ): + super().__init__() + channels = tuple(channels) + self.in_channels = in_channels + self.out_channels = out_channels + self.time_embeddings = SinusoidalPosEmb(in_channels) + time_embed_dim = channels[0] * 4 + self.time_mlp = TimestepEmbedding( + in_channels=in_channels, time_embed_dim=time_embed_dim, act_fn="silu" + ) + self.down_blocks = paddle.nn.LayerList(sublayers=[]) + self.mid_blocks = paddle.nn.LayerList(sublayers=[]) + self.up_blocks = paddle.nn.LayerList(sublayers=[]) + output_channel = in_channels + for i in range(len(channels)): + input_channel = output_channel + output_channel = channels[i] + is_last = i == len(channels) - 1 + resnet = ResnetBlock1D( + dim=input_channel, dim_out=output_channel, time_emb_dim=time_embed_dim + ) + transformer_blocks = paddle.nn.LayerList( + sublayers=[ + self.get_block( + down_block_type, + output_channel, + attention_head_dim, + num_heads, + dropout, + act_fn, + ) + for _ in range(n_blocks) + ] + ) + downsample = ( + Downsample1D(output_channel) + if not is_last + else paddle.nn.Conv1d(output_channel, output_channel, 3, padding=1) + ) + self.down_blocks.append( + paddle.nn.LayerList(sublayers=[resnet, transformer_blocks, downsample]) + ) + for i in range(num_mid_blocks): + input_channel = channels[-1] + out_channels = channels[-1] + resnet = ResnetBlock1D( + dim=input_channel, dim_out=output_channel, time_emb_dim=time_embed_dim + ) + transformer_blocks = paddle.nn.LayerList( + sublayers=[ + self.get_block( + mid_block_type, + output_channel, + attention_head_dim, + num_heads, + dropout, + act_fn, + ) + for _ in range(n_blocks) + ] + ) + self.mid_blocks.append( + paddle.nn.LayerList(sublayers=[resnet, transformer_blocks]) + ) + channels = channels[::-1] + (channels[0],) + for i in range(len(channels) - 1): + input_channel = channels[i] + output_channel = channels[i + 1] + is_last = i == len(channels) - 2 + resnet = ResnetBlock1D( + dim=2 * input_channel, + dim_out=output_channel, + time_emb_dim=time_embed_dim, + ) + transformer_blocks = paddle.nn.LayerList( + sublayers=[ + self.get_block( + up_block_type, + output_channel, + attention_head_dim, + num_heads, + dropout, + act_fn, + ) + for _ in range(n_blocks) + ] + ) + upsample = ( + Upsample1D(output_channel, use_conv_transpose=True) + if not is_last + else paddle.nn.Conv1d(output_channel, output_channel, 3, padding=1) + ) + self.up_blocks.append( + paddle.nn.LayerList(sublayers=[resnet, transformer_blocks, upsample]) + ) + self.final_block = Block1D(channels[-1], channels[-1]) + self.final_proj = paddle.nn.Conv1d(channels[-1], self.out_channels, 1) + self.initialize_weights() + + @staticmethod + def get_block(block_type, dim, attention_head_dim, num_heads, dropout, act_fn): + if block_type == "conformer": + block = ConformerWrapper( + dim=dim, + dim_head=attention_head_dim, + heads=num_heads, + ff_mult=1, + conv_expansion_factor=2, + ff_dropout=dropout, + attn_dropout=dropout, + conv_dropout=dropout, + conv_kernel_size=31, + ) + elif block_type == "transformer": + block = BasicTransformerBlock( + dim=dim, + num_attention_heads=num_heads, + attention_head_dim=attention_head_dim, + dropout=dropout, + activation_fn=act_fn, + ) + else: + raise ValueError(f"Unknown block type {block_type}") + return block + + def initialize_weights(self): + for m in self.sublayers(): + if isinstance(m, paddle.nn.Conv1d): + paddle.nn.init.kaiming_normal_(m.weight, nonlinearity="relu") + if m.bias is not None: + paddle.nn.init.constant_(m.bias, 0) + elif isinstance(m, paddle.nn.GroupNorm): + paddle.nn.init.constant_(m.weight, 1) + paddle.nn.init.constant_(m.bias, 0) + elif isinstance(m, paddle.nn.Linear): + paddle.nn.init.kaiming_normal_(m.weight, nonlinearity="relu") + if m.bias is not None: + paddle.nn.init.constant_(m.bias, 0) + + def forward(self, x, mask, mu, t, spks=None, cond=None): + """Forward pass of the UNet1DConditional model. + + Args: + x (torch.Tensor): shape (batch_size, in_channels, time) + mask (_type_): shape (batch_size, 1, time) + t (_type_): shape (batch_size) + spks (_type_, optional): shape: (batch_size, condition_channels). Defaults to None. + cond (_type_, optional): placeholder for future use. Defaults to None. + + Raises: + ValueError: _description_ + ValueError: _description_ + + Returns: + _type_: _description_ + """ + t = self.time_embeddings(t) + t = self.time_mlp(t) + x = einops.pack([x, mu], "b * t")[0] + if spks is not None: + spks = einops.repeat(spks, "b c -> b c t", t=x.shape[-1]) + x = einops.pack([x, spks], "b * t")[0] + hiddens = [] + masks = [mask] + for resnet, transformer_blocks, downsample in self.down_blocks: + mask_down = masks[-1] + x = resnet(x, mask_down, t) + x = einops.rearrange(x, "b c t -> b t c") + mask_down = einops.rearrange(mask_down, "b 1 t -> b t") + for transformer_block in transformer_blocks: + x = transformer_block( + hidden_states=x, attention_mask=mask_down, timestep=t + ) + x = einops.rearrange(x, "b t c -> b c t") + mask_down = einops.rearrange(mask_down, "b t -> b 1 t") + hiddens.append(x) + x = downsample(x * mask_down) + masks.append(mask_down[:, :, ::2]) + masks = masks[:-1] + mask_mid = masks[-1] + for resnet, transformer_blocks in self.mid_blocks: + x = resnet(x, mask_mid, t) + x = einops.rearrange(x, "b c t -> b t c") + mask_mid = einops.rearrange(mask_mid, "b 1 t -> b t") + for transformer_block in transformer_blocks: + x = transformer_block( + hidden_states=x, attention_mask=mask_mid, timestep=t + ) + x = einops.rearrange(x, "b t c -> b c t") + mask_mid = einops.rearrange(mask_mid, "b t -> b 1 t") + for resnet, transformer_blocks, upsample in self.up_blocks: + mask_up = masks.pop() + x = resnet(einops.pack([x, hiddens.pop()], "b * t")[0], mask_up, t) + x = einops.rearrange(x, "b c t -> b t c") + mask_up = einops.rearrange(mask_up, "b 1 t -> b t") + for transformer_block in transformer_blocks: + x = transformer_block( + hidden_states=x, attention_mask=mask_up, timestep=t + ) + x = einops.rearrange(x, "b t c -> b c t") + mask_up = einops.rearrange(mask_up, "b t -> b 1 t") + x = upsample(x * mask_up) + x = self.final_block(x, mask_up) + output = self.final_proj(x * mask_up) + return output * mask diff --git a/paddlespeech/t2s/modules/flow/matcha_transformer.py b/paddlespeech/t2s/modules/flow/matcha_transformer.py new file mode 100644 index 000000000..3be13dd10 --- /dev/null +++ b/paddlespeech/t2s/modules/flow/matcha_transformer.py @@ -0,0 +1,318 @@ +from typing import Any, Dict, Optional +from .attention_processor import Attention +import paddle +from paddle import nn +from paddlespeech.t2s.modules.flow.lora import LoRACompatibleLinear +class SnakeBeta(paddle.nn.Layer): + """ + A modified Snake function which uses separate parameters for the magnitude of the periodic components + Shape: + - Input: (B, C, T) + - Output: (B, C, T), same shape as the input + Parameters: + - alpha - trainable parameter that controls frequency + - beta - trainable parameter that controls magnitude + References: + - This activation function is a modified version based on this paper by Liu Ziyin, Tilman Hartwig, Masahito Ueda: + https://arxiv.org/abs/2006.08195 + Examples: + >>> a1 = snakebeta(256) + >>> x = torch.randn(256) + >>> x = a1(x) + """ + + def __init__( + self, + in_features, + out_features, + alpha=1.0, + alpha_trainable=True, + alpha_logscale=True, + ): + """ + Initialization. + INPUT: + - in_features: shape of the input + - alpha - trainable parameter that controls frequency + - beta - trainable parameter that controls magnitude + alpha is initialized to 1 by default, higher values = higher-frequency. + beta is initialized to 1 by default, higher values = higher-magnitude. + alpha will be trained along with the rest of your model. + """ + super().__init__() + self.in_features = ( + out_features if isinstance(out_features, list) else [out_features] + ) + self.proj = LoRACompatibleLinear( + in_features, out_features + ) + self.alpha_logscale = alpha_logscale + if self.alpha_logscale: + self.alpha = paddle.nn.parameter.Parameter( + paddle.zeros(self.in_features) * alpha + ) + self.beta = paddle.nn.parameter.Parameter( + paddle.zeros(self.in_features) * alpha + ) + else: + self.alpha = paddle.nn.parameter.Parameter( + paddle.ones(self.in_features) * alpha + ) + self.beta = paddle.nn.parameter.Parameter( + paddle.ones(self.in_features) * alpha + ) + self.alpha.stop_gradient = not alpha_trainable + self.beta.stop_gradient = not alpha_trainable + self.no_div_by_zero = 1e-09 + + def forward(self, x): + """ + Forward pass of the function. + Applies the function to the input elementwise. + SnakeBeta ∶= x + 1/b * sin^2 (xa) + """ + x = self.proj(x) + if self.alpha_logscale: + alpha = paddle.exp(x=self.alpha) + beta = paddle.exp(x=self.beta) + else: + alpha = self.alpha + beta = self.beta + x = x + 1.0 / (beta + self.no_div_by_zero) * paddle.pow( + paddle.sin(x * alpha), 2 + ) + return x +import paddle +import paddle.nn as nn +import paddle.nn.functional as F + +class GELU(nn.Layer): + def __init__(self, dim_in: int, dim_out: int, approximate: str = "none", bias: bool = True): + super().__init__() + self.proj = nn.Linear(dim_in, dim_out, bias_attr=bias) + self.approximate = approximate + + def gelu(self, gate: paddle.Tensor) -> paddle.Tensor: + if self.approximate == "tanh": + approximate_bool = True + else: + approximate_bool = False + + if gate.dtype == paddle.float16: + return F.gelu(gate.astype(paddle.float32), approximate=approximate_bool).astype(paddle.float16) + else: + return F.gelu(gate, approximate=approximate_bool) + + def forward(self, hidden_states: paddle.Tensor) -> paddle.Tensor: + hidden_states = self.proj(hidden_states) + hidden_states = self.gelu(hidden_states) + return hidden_states + +class FeedForward(paddle.nn.Layer): + """ + A feed-forward layer. + + Parameters: + dim (`int`): The number of channels in the input. + dim_out (`int`, *optional*): The number of channels in the output. If not given, defaults to `dim`. + mult (`int`, *optional*, defaults to 4): The multiplier to use for the hidden dimension. + dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use. + activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward. + final_dropout (`bool` *optional*, defaults to False): Apply a final dropout. + """ + + def __init__( + self, + dim: int, + dim_out: Optional[int] = None, + mult: int = 4, + dropout: float = 0.0, + activation_fn: str = "geglu", + final_dropout: bool = False, + ): + super().__init__() + inner_dim = int(dim * mult) + dim_out = dim_out if dim_out is not None else dim + act_fn = GELU(dim, inner_dim) +# if activation_fn == "gelu": +# >>>>>> act_fn = diffusers.models.attention.GELU(dim, inner_dim) +# if activation_fn == "gelu-approximate": +# >>>>>> act_fn = diffusers.models.attention.GELU(dim, inner_dim, approximate="tanh") +# elif activation_fn == "geglu": +# >>>>>> act_fn = diffusers.models.attention.GEGLU(dim, inner_dim) +# elif activation_fn == "geglu-approximate": +# act_fn = diffusers.models.attention.ApproximateGELU(dim, inner_dim) +# elif activation_fn == "snakebeta": + # act_fn = SnakeBeta(dim, inner_dim) + self.net = paddle.nn.LayerList(sublayers=[]) + self.net.append(act_fn) + self.net.append(paddle.nn.Dropout(p=dropout)) + self.net.append(LoRACompatibleLinear(inner_dim, dim_out)) + if final_dropout: + self.net.append(paddle.nn.Dropout(p=dropout)) + + def forward(self, hidden_states): + for module in self.net: + hidden_states = module(hidden_states) + return hidden_states + + + +class BasicTransformerBlock(paddle.nn.Layer): + """ + A basic Transformer block. + + Parameters: + dim (`int`): The number of channels in the input and output. + num_attention_heads (`int`): The number of heads to use for multi-head attention. + attention_head_dim (`int`): The number of channels in each head. + dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use. + cross_attention_dim (`int`, *optional*): The size of the encoder_hidden_states vector for cross attention. + only_cross_attention (`bool`, *optional*): + Whether to use only cross-attention layers. In this case two cross attention layers are used. + double_self_attention (`bool`, *optional*): + Whether to use two self-attention layers. In this case no cross attention layers are used. + activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward. + num_embeds_ada_norm (: + obj: `int`, *optional*): The number of diffusion steps used during training. See `Transformer2DModel`. + attention_bias (: + obj: `bool`, *optional*, defaults to `False`): Configure if the attentions should contain a bias parameter. + """ + + def __init__( + self, + dim: int, + num_attention_heads: int, + attention_head_dim: int, + dropout=0.0, + cross_attention_dim: Optional[int] = None, + activation_fn: str = "geglu", + num_embeds_ada_norm: Optional[int] = None, + attention_bias: bool = False, + only_cross_attention: bool = False, + double_self_attention: bool = False, + upcast_attention: bool = False, + norm_elementwise_affine: bool = True, + norm_type: str = "layer_norm", + final_dropout: bool = False, + ): + super().__init__() + self.only_cross_attention = only_cross_attention + self.use_ada_layer_norm_zero = ( + num_embeds_ada_norm is not None and norm_type == "ada_norm_zero" + ) + self.use_ada_layer_norm = ( + num_embeds_ada_norm is not None and norm_type == "ada_norm" + ) + if norm_type in ("ada_norm", "ada_norm_zero") and num_embeds_ada_norm is None: + raise ValueError( + f"`norm_type` is set to {norm_type}, but `num_embeds_ada_norm` is not defined. Please make sure to define `num_embeds_ada_norm` if setting `norm_type` to {norm_type}." + ) + + self.norm1 = paddle.nn.LayerNorm( + normalized_shape=dim, + weight_attr=norm_elementwise_affine, + bias_attr=norm_elementwise_affine, + ) + self.attn1 = Attention( + query_dim=dim, + heads=num_attention_heads, + dim_head=attention_head_dim, + dropout=dropout, + bias=attention_bias, + cross_attention_dim=cross_attention_dim if only_cross_attention else None, + upcast_attention=upcast_attention, + ) + self.norm2 = None + self.attn2 = None + self.norm3 = paddle.nn.LayerNorm( + normalized_shape=dim, + weight_attr=norm_elementwise_affine, + bias_attr=norm_elementwise_affine, + ) + self.ff = FeedForward( + dim, + dropout=dropout, + activation_fn=activation_fn, + final_dropout=final_dropout, + ) + self._chunk_size = None + self._chunk_dim = 0 + + def set_chunk_feed_forward(self, chunk_size: Optional[int], dim: int): + self._chunk_size = chunk_size + self._chunk_dim = dim + + def forward( + self, + hidden_states: paddle.Tensor, + attention_mask: Optional[paddle.Tensor] = None, + encoder_hidden_states: Optional[paddle.Tensor] = None, + encoder_attention_mask: Optional[paddle.Tensor] = None, + timestep: Optional[paddle.Tensor] = None, + cross_attention_kwargs: Dict[str, Any] = None, + class_labels: Optional[paddle.Tensor] = None, + ): + if self.use_ada_layer_norm: + norm_hidden_states = self.norm1(hidden_states, timestep) + elif self.use_ada_layer_norm_zero: + (norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp) = self.norm1( + hidden_states, timestep, class_labels, hidden_dtype=hidden_states.dtype + ) + else: + norm_hidden_states = self.norm1(hidden_states) + cross_attention_kwargs = ( + cross_attention_kwargs if cross_attention_kwargs is not None else {} + ) + attn_output = self.attn1( + norm_hidden_states, + encoder_hidden_states=encoder_hidden_states + if self.only_cross_attention + else None, + attention_mask=encoder_attention_mask + if self.only_cross_attention + else attention_mask, + **cross_attention_kwargs, + ) + if self.use_ada_layer_norm_zero: + attn_output = gate_msa.unsqueeze(1) * attn_output + hidden_states = attn_output + hidden_states + if self.attn2 is not None: + norm_hidden_states = ( + self.norm2(hidden_states, timestep) + if self.use_ada_layer_norm + else self.norm2(hidden_states) + ) + attn_output = self.attn2( + norm_hidden_states, + encoder_hidden_states=encoder_hidden_states, + attention_mask=encoder_attention_mask, + **cross_attention_kwargs, + ) + hidden_states = attn_output + hidden_states + norm_hidden_states = self.norm3(hidden_states) + if self.use_ada_layer_norm_zero: + norm_hidden_states = ( + norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] + ) + if self._chunk_size is not None: + if norm_hidden_states.shape[self._chunk_dim] % self._chunk_size != 0: + raise ValueError( + f"`hidden_states` dimension to be chunked: {norm_hidden_states.shape[self._chunk_dim]} has to be divisible by chunk size: {self._chunk_size}. Make sure to set an appropriate `chunk_size` when calling `unet.enable_forward_chunking`." + ) + num_chunks = norm_hidden_states.shape[self._chunk_dim] // self._chunk_size + ff_output = paddle.cat( + [ + self.ff(hid_slice) + for hid_slice in norm_hidden_states.chunk( + num_chunks, dim=self._chunk_dim + ) + ], + dim=self._chunk_dim, + ) + else: + ff_output = self.ff(norm_hidden_states) + if self.use_ada_layer_norm_zero: + ff_output = gate_mlp.unsqueeze(1) * ff_output + hidden_states = ff_output + hidden_states + return hidden_states diff --git a/paddlespeech/t2s/modules/flow/normalization.py b/paddlespeech/t2s/modules/flow/normalization.py new file mode 100644 index 000000000..ba0ccbc55 --- /dev/null +++ b/paddlespeech/t2s/modules/flow/normalization.py @@ -0,0 +1,276 @@ +import numbers +from typing import Dict, Optional, Tuple + +import paddle +from paddle import nn +from .activations import get_activation +from .embeddings import (CombinedTimestepLabelEmbeddings, + PixArtAlphaCombinedTimestepSizeEmbeddings) +def get_activation(act_fn): + if act_fn == "silu": + return nn.Silu() + elif act_fn == "mish": + return nn.Mish() + elif act_fn == "relu": + return nn.ReLU() + elif act_fn == "gelu": + return nn.GELU() + else: + raise ValueError(f"Unsupported activation function: {act_fn}") + +class AdaLayerNorm(paddle.nn.Layer): + """ + Norm layer modified to incorporate timestep embeddings. + + Parameters: + embedding_dim (`int`): The size of each embedding vector. + num_embeddings (`int`): The size of the embeddings dictionary. + """ + + def __init__(self, embedding_dim: int, num_embeddings: int): + super().__init__() + self.emb = paddle.nn.Embedding(num_embeddings, embedding_dim) + self.silu = paddle.nn.SiLU() + self.linear = paddle.nn.Linear( + in_features=embedding_dim, out_features=embedding_dim * 2 + ) + self.norm = paddle.nn.LayerNorm( + normalized_shape=embedding_dim, weight_attr=False, bias_attr=False + ) + + def forward(self, x: paddle.Tensor, timestep: paddle.Tensor) -> paddle.Tensor: + emb = self.linear(self.silu(self.emb(timestep))) + scale, shift = paddle.chunk(emb, 2) + x = self.norm(x) * (1 + scale) + shift + return x + + +class AdaLayerNormZero(paddle.nn.Layer): + """ + Norm layer adaptive layer norm zero (adaLN-Zero). + + Parameters: + embedding_dim (`int`): The size of each embedding vector. + num_embeddings (`int`): The size of the embeddings dictionary. + """ + + def __init__(self, embedding_dim: int, num_embeddings: Optional[int] = None): + super().__init__() + if num_embeddings is not None: + self.emb = CombinedTimestepLabelEmbeddings(num_embeddings, embedding_dim) + else: + self.emb = None + self.silu = paddle.nn.SiLU() + self.linear = paddle.nn.Linear( + in_features=embedding_dim, out_features=6 * embedding_dim, bias_attr=True + ) + self.norm = paddle.nn.LayerNorm( + normalized_shape=embedding_dim, + weight_attr=False, + bias_attr=False, + epsilon=1e-06, + ) + + def forward( + self, + x: paddle.Tensor, + timestep: Optional[paddle.Tensor] = None, + class_labels: Optional[paddle.LongTensor] = None, + hidden_dtype: Optional[paddle.dtype] = None, + emb: Optional[paddle.Tensor] = None, + ) -> Tuple[ + paddle.Tensor, paddle.Tensor, paddle.Tensor, paddle.Tensor, paddle.Tensor + ]: + if self.emb is not None: + emb = self.emb(timestep, class_labels, hidden_dtype=hidden_dtype) + emb = self.linear(self.silu(emb)) + shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = emb.chunk( + 6, dim=1 + ) + x = self.norm(x) * (1 + scale_msa[:, None]) + shift_msa[:, None] + return x, gate_msa, shift_mlp, scale_mlp, gate_mlp + + +class AdaLayerNormSingle(paddle.nn.Layer): + """ + Norm layer adaptive layer norm single (adaLN-single). + + As proposed in PixArt-Alpha (see: https://arxiv.org/abs/2310.00426; Section 2.3). + + Parameters: + embedding_dim (`int`): The size of each embedding vector. + use_additional_conditions (`bool`): To use additional conditions for normalization or not. + """ + + def __init__(self, embedding_dim: int, use_additional_conditions: bool = False): + super().__init__() + self.emb = PixArtAlphaCombinedTimestepSizeEmbeddings( + embedding_dim, + size_emb_dim=embedding_dim // 3, + use_additional_conditions=use_additional_conditions, + ) + self.silu = paddle.nn.SiLU() + self.linear = paddle.nn.Linear( + in_features=embedding_dim, out_features=6 * embedding_dim, bias_attr=True + ) + + def forward( + self, + timestep: paddle.Tensor, + added_cond_kwargs: Optional[Dict[str, paddle.Tensor]] = None, + batch_size: Optional[int] = None, + hidden_dtype: Optional[paddle.dtype] = None, + ) -> Tuple[ + paddle.Tensor, paddle.Tensor, paddle.Tensor, paddle.Tensor, paddle.Tensor + ]: + embedded_timestep = self.emb( + timestep, + **added_cond_kwargs, + batch_size=batch_size, + hidden_dtype=hidden_dtype, + ) + return self.linear(self.silu(embedded_timestep)), embedded_timestep + + +class AdaGroupNorm(paddle.nn.Layer): + """ + GroupNorm layer modified to incorporate timestep embeddings. + + Parameters: + embedding_dim (`int`): The size of each embedding vector. + num_embeddings (`int`): The size of the embeddings dictionary. + num_groups (`int`): The number of groups to separate the channels into. + act_fn (`str`, *optional*, defaults to `None`): The activation function to use. + eps (`float`, *optional*, defaults to `1e-5`): The epsilon value to use for numerical stability. + """ + + def __init__( + self, + embedding_dim: int, + out_dim: int, + num_groups: int, + act_fn: Optional[str] = None, + eps: float = 1e-05, + ): + super().__init__() + self.num_groups = num_groups + self.eps = eps + if act_fn is None: + self.act = None + else: + self.act = get_activation(act_fn) + self.linear = paddle.nn.Linear( + in_features=embedding_dim, out_features=out_dim * 2 + ) + + def forward(self, x: paddle.Tensor, emb: paddle.Tensor) -> paddle.Tensor: + if self.act: + emb = self.act(emb) + emb = self.linear(emb) + emb = emb[:, :, None, None] + scale, shift = emb.chunk(2, dim=1) + x = paddle.nn.functional.group_norm( + x=x, num_groups=self.num_groups, epsilon=self.eps + ) + x = x * (1 + scale) + shift + return x + + +class AdaLayerNormContinuous(paddle.nn.Layer): + def __init__( + self, + embedding_dim: int, + conditioning_embedding_dim: int, + elementwise_affine=True, + eps=1e-05, + bias=True, + norm_type="layer_norm", + ): + super().__init__() + self.silu = paddle.nn.SiLU() + self.linear = paddle.nn.Linear( + in_features=conditioning_embedding_dim, + out_features=embedding_dim * 2, + bias_attr=bias, + ) + if norm_type == "layer_norm": + self.norm = LayerNorm(embedding_dim, eps, elementwise_affine, bias) + elif norm_type == "rms_norm": + self.norm = RMSNorm(embedding_dim, eps, elementwise_affine) + else: + raise ValueError(f"unknown norm_type {norm_type}") + + def forward( + self, x: paddle.Tensor, conditioning_embedding: paddle.Tensor + ) -> paddle.Tensor: + emb = self.linear(self.silu(conditioning_embedding).to(x.dtype)) + scale, shift = paddle.chunk(emb, 2, dim=1) + x = self.norm(x) * (1 + scale)[:, None, :] + shift[:, None, :] + return x + +LayerNorm = paddle.nn.LayerNorm + + +class LayerNorm(paddle.nn.Layer): + def __init__( + self, + dim, + eps: float = 1e-05, + elementwise_affine: bool = True, + bias: bool = True, + ): + super().__init__() + self.eps = eps + if isinstance(dim, numbers.Integral): + dim = (dim,) + self.dim = paddle.Size(dim) + if elementwise_affine: + self.weight = paddle.nn.parameter.Parameter(paddle.ones(dim)) + self.bias = ( + paddle.nn.parameter.Parameter(paddle.zeros(dim)) if bias else None + ) + else: + self.weight = None + self.bias = None + + def forward(self, input): + return paddle.nn.functional.layer_norm( + input, self.dim, self.weight, self.bias, self.eps + ) + + +class RMSNorm(paddle.nn.Layer): + def __init__(self, dim, eps: float, elementwise_affine: bool = True): + super().__init__() + self.eps = eps + if isinstance(dim, numbers.Integral): + dim = (dim,) + self.dim = paddle.Size(dim) + if elementwise_affine: + self.weight = paddle.nn.parameter.Parameter(paddle.ones(dim)) + else: + self.weight = None + + def forward(self, hidden_states): + input_dtype = hidden_states.dtype + variance = hidden_states.to(paddle.float32).pow(2).mean(-1, keepdim=True) + hidden_states = hidden_states * paddle.rsqrt(variance + self.eps) + if self.weight is not None: + if self.weight.dtype in [paddle.float16, paddle.bfloat16]: + hidden_states = hidden_states.to(self.weight.dtype) + hidden_states = hidden_states * self.weight + else: + hidden_states = hidden_states.to(input_dtype) + return hidden_states + + +class GlobalResponseNorm(paddle.nn.Layer): + def __init__(self, dim): + super().__init__() + self.gamma = paddle.nn.parameter.Parameter(paddle.zeros(1, 1, 1, dim)) + self.beta = paddle.nn.parameter.Parameter(paddle.zeros(1, 1, 1, dim)) + + def forward(self, x): + gx = paddle.norm(x, p=2, dim=(1, 2), keepdim=True) + nx = gx / (gx.mean(dim=-1, keepdim=True) + 1e-06) + return self.gamma * (x * nx) + self.beta + x diff --git a/paddlespeech/t2s/modules/predictor/length_regulator.py b/paddlespeech/t2s/modules/predictor/length_regulator.py index bdfa18391..bb4b5a64d 100644 --- a/paddlespeech/t2s/modules/predictor/length_regulator.py +++ b/paddlespeech/t2s/modules/predictor/length_regulator.py @@ -108,7 +108,6 @@ class LengthRegulator(nn.Layer): Returns: Tensor: replicated input tensor based on durations (B, T*, D). """ - if alpha != 1.0: assert alpha > 0 ds = paddle.round(ds.cast(dtype=paddle.float32) * alpha) diff --git a/paddlespeech/t2s/modules/transformer/activation.py b/paddlespeech/t2s/modules/transformer/activation.py new file mode 100644 index 000000000..5380078ac --- /dev/null +++ b/paddlespeech/t2s/modules/transformer/activation.py @@ -0,0 +1,82 @@ +import paddle + +"""Swish() activation function for Conformer.""" + + +class Swish(paddle.nn.Layer): + """Construct an Swish object.""" + + def forward(self, x: paddle.Tensor) -> paddle.Tensor: + """Return Swish activation function.""" + return x * paddle.nn.functional.sigmoid(x) + + +class Snake(paddle.nn.Layer): + ''' + Implementation of a sine-based periodic activation function + Shape: + - Input: (B, C, T) + - Output: (B, C, T), same shape as the input + Parameters: + - alpha - trainable parameter + References: + - This activation function is from this paper by Liu Ziyin, Tilman Hartwig, Masahito Ueda: + https://arxiv.org/abs/2006.08195 + Examples: + >>> a1 = Snake(256) + >>> x = paddle.randn([1, 256, 100]) # Example input + >>> x = a1(x) + ''' + def __init__(self, in_features, alpha=1.0, alpha_trainable=True, alpha_logscale=False): + ''' + Initialization. + INPUT: + - in_features: number of input features (channel dimension) + - alpha: trainable parameter + alpha is initialized to 1 by default, higher values = higher-frequency. + alpha will be trained along with the rest of your model. + ''' + super(Snake, self).__init__() + self.in_features = in_features + self.alpha_logscale = alpha_logscale + + # 避免除零的小常数 + self.no_div_by_zero = 1e-9 + + # 初始化alpha的值:对数尺度下初始为0,线性尺度下初始为alpha + if self.alpha_logscale: + initial_value = 0.0 # 对数尺度下,初始化为0,前向传播中会进行exp运算 + else: + initial_value = alpha # 线性尺度下,直接初始化为alpha + + # 创建可训练参数alpha - 使用PaddlePaddle的方式 + # 注意:这里使用self.create_parameter而不是paddle.create_parameter + self.alpha = self.create_parameter( + shape=[in_features], # 参数形状为[in_features] + dtype='float32', # 数据类型 + default_initializer=paddle.nn.initializer.Constant(value=initial_value) # 初始化器 + ) + + # 设置参数是否需要梯度更新(是否可训练) + # 在PaddlePaddle中,通过设置stop_gradient来控制 + self.alpha.stop_gradient = not alpha_trainable + + def forward(self, x): + ''' + Forward pass of the function. + Applies the function to the input elementwise. + Snake ∶= x + 1/a * sin^2 (xa) + ''' + # 调整alpha的维度以匹配输入x: [B, C, T] -> alpha需要变为[1, C, 1] + alpha = self.alpha.unsqueeze(0).unsqueeze(-1) # 从[C]变为[1, C, 1] + + # 如果使用对数尺度,对alpha取指数 + if self.alpha_logscale: + alpha = paddle.exp(alpha) + + # 计算Snake激活函数 + # 公式: x + (1.0 / (alpha + epsilon)) * sin(x * alpha)^2 + sin_term = paddle.sin(x * alpha) + result = x + (1.0 / (alpha + self.no_div_by_zero)) * (sin_term ** 2) + + return result diff --git a/paddlespeech/t2s/modules/transformer/attention.py b/paddlespeech/t2s/modules/transformer/attention.py index 3237be1b6..0ef07c252 100644 --- a/paddlespeech/t2s/modules/transformer/attention.py +++ b/paddlespeech/t2s/modules/transformer/attention.py @@ -199,7 +199,7 @@ class RelPositionMultiHeadedAttention(MultiHeadedAttention): x = x * paddle.tril(ones, t2 - t1)[None, None, :, :] return x - def forward(self, query, key, value, pos_emb, mask): + def forward(self, query, key, value, pos_emb, mask, cache): """Compute 'Scaled Dot Product Attention' with rel. positional encoding. Args: @@ -220,6 +220,11 @@ class RelPositionMultiHeadedAttention(MultiHeadedAttention): q, k, v = self.forward_qkv(query, key, value) # (batch, time1, head, d_k) q = q.transpose([0, 2, 1, 3]) + if cache is not None and cache.shape[0] > 0: + key_cache, value_cache = paddle.split(cache, num_or_sections=2, axis=-1) + k = paddle.concat([key_cache, k], axis=2) + v = paddle.concat([value_cache, v], axis=2) + new_cache = paddle.concat([k, v], axis=-1) n_batch_pos = paddle.shape(pos_emb)[0] p = self.linear_pos(pos_emb).reshape( [n_batch_pos, -1, self.h, self.d_k]) @@ -243,7 +248,7 @@ class RelPositionMultiHeadedAttention(MultiHeadedAttention): # (batch, head, time1, time2) scores = (matrix_ac + matrix_bd) / math.sqrt(self.d_k) - return self.forward_attention(v, scores, mask) + return self.forward_attention(v, scores, mask), new_cache class LegacyRelPositionMultiHeadedAttention(MultiHeadedAttention): diff --git a/paddlespeech/t2s/modules/transformer/convolution.py b/paddlespeech/t2s/modules/transformer/convolution.py new file mode 100644 index 000000000..9f38479d8 --- /dev/null +++ b/paddlespeech/t2s/modules/transformer/convolution.py @@ -0,0 +1,99 @@ +import paddle + +"""ConvolutionModule definition.""" +from typing import Tuple + + +class ConvolutionModule(paddle.nn.Layer): + """ConvolutionModule in Conformer model.""" + + def __init__( + self, + channels: int, + kernel_size: int = 15, + activation: paddle.nn.Layer = paddle.nn.ReLU(), + norm: str = "batch_norm", + causal: bool = False, + bias: bool = True, + ): + """Construct an ConvolutionModule object. + Args: + channels (int): The number of channels of conv layers. + kernel_size (int): Kernel size of conv layers. + causal (int): Whether use causal convolution or not + """ + super().__init__() + self.pointwise_conv1 = paddle.nn.Conv1d( + channels, 2 * channels, kernel_size=1, stride=1, padding=0, bias=bias + ) + if causal: + padding = 0 + self.lorder = kernel_size - 1 + else: + assert (kernel_size - 1) % 2 == 0 + padding = (kernel_size - 1) // 2 + self.lorder = 0 + self.depthwise_conv = paddle.nn.Conv1d( + channels, + channels, + kernel_size, + stride=1, + padding=padding, + groups=channels, + bias=bias, + ) + assert norm in ["batch_norm", "layer_norm"] + if norm == "batch_norm": + self.use_layer_norm = False + self.norm = paddle.nn.BatchNorm1D(num_features=channels) + else: + self.use_layer_norm = True + self.norm = paddle.nn.LayerNorm(normalized_shape=channels) + self.pointwise_conv2 = paddle.nn.Conv1d( + channels, channels, kernel_size=1, stride=1, padding=0, bias=bias + ) + self.activation = activation + + def forward( + self, + x: paddle.Tensor, + mask_pad: paddle.Tensor = paddle.ones((0, 0, 0), dtype=paddle.bool), + cache: paddle.Tensor = paddle.zeros((0, 0, 0)), + ) -> Tuple[paddle.Tensor, paddle.Tensor]: + """Compute convolution module. + Args: + x (torch.Tensor): Input tensor (#batch, time, channels). + mask_pad (torch.Tensor): used for batch padding (#batch, 1, time), + (0, 0, 0) means fake mask. + cache (torch.Tensor): left context cache, it is only + used in causal convolution (#batch, channels, cache_t), + (0, 0, 0) meas fake cache. + Returns: + torch.Tensor: Output tensor (#batch, time, channels). + """ + x = x.transpose(1, 2) + if mask_pad.size(2) > 0: + x.masked_fill_(~mask_pad, 0.0) + if self.lorder > 0: + if cache.size(2) == 0: + x = paddle.compat.pad(x, (self.lorder, 0), "constant", 0.0) + else: + assert cache.size(0) == x.size(0) + assert cache.size(1) == x.size(1) + x = paddle.cat((cache, x), dim=2) + assert x.size(2) > self.lorder + new_cache = x[:, :, -self.lorder :] + else: + new_cache = paddle.zeros((0, 0, 0), dtype=x.dtype, device=x.place) + x = self.pointwise_conv1(x) + x = paddle.nn.functional.glu(x=x, axis=1) + x = self.depthwise_conv(x) + if self.use_layer_norm: + x = x.transpose(1, 2) + x = self.activation(self.norm(x)) + if self.use_layer_norm: + x = x.transpose(1, 2) + x = self.pointwise_conv2(x) + if mask_pad.size(2) > 0: + x.masked_fill_(~mask_pad, 0.0) + return x.transpose(1, 2), new_cache diff --git a/paddlespeech/t2s/modules/transformer/embedding.py b/paddlespeech/t2s/modules/transformer/embedding.py index e4331cff0..24a076f31 100644 --- a/paddlespeech/t2s/modules/transformer/embedding.py +++ b/paddlespeech/t2s/modules/transformer/embedding.py @@ -14,7 +14,7 @@ # Modified from espnet(https://github.com/espnet/espnet) """Positional Encoding Module.""" import math - +from typing import Union import paddle from paddle import nn @@ -131,6 +131,108 @@ class ScaledPositionalEncoding(PositionalEncoding): x = x + self.alpha * self.pe[:, :T] return self.dropout(x) +class EspnetRelPositionalEncoding(paddle.nn.Layer): + """Relative positional encoding module (new implementation). + + Details can be found in https://github.com/espnet/espnet/pull/2816. + + See : Appendix B in https://arxiv.org/abs/1901.02860 + + Args: + d_model (int): Embedding dimension. + dropout_rate (float): Dropout rate. + max_len (int): Maximum input length. + + """ + + def __init__(self, d_model: int, dropout_rate: float, max_len: int = 5000): + """Construct an PositionalEncoding object.""" + super(EspnetRelPositionalEncoding, self).__init__() + self.d_model = d_model + self.xscale = math.sqrt(self.d_model) + self.dropout = paddle.nn.Dropout(p=dropout_rate) + self.pe = None + self.extend_pe(paddle.to_tensor([0.0]).expand([1, max_len])) + + def extend_pe(self, x: paddle.Tensor): + """Reset the positional encodings.""" + if self.pe is not None: + if self.pe.shape[1] >= x.shape[1] * 2 - 1: + if self.pe.dtype != x.dtype or self.pe.place != x.place: + self.pe = self.pe.to(dtype=x.dtype, device=x.place) + return + pe_positive = paddle.zeros([x.shape[1], self.d_model]) + pe_negative = paddle.zeros([x.shape[1], self.d_model]) + position = paddle.arange(0, x.shape[1], dtype=paddle.float32).unsqueeze(1) + div_term = paddle.exp( + x=paddle.arange(0, self.d_model, 2, dtype=paddle.float32) + * -(math.log(10000.0) / self.d_model) + ) + pe_positive[:, 0::2] = paddle.sin(position * div_term) + pe_positive[:, 1::2] = paddle.cos(position * div_term) + pe_negative[:, 0::2] = paddle.sin(-1 * position * div_term) + pe_negative[:, 1::2] = paddle.cos(-1 * position * div_term) + pe_positive = paddle.flip(x=pe_positive, axis=[0]).unsqueeze(0) + pe_negative = pe_negative[1:].unsqueeze(0) + pe = paddle.cat([pe_positive, pe_negative], dim=1) + self.pe = pe.to(device=x.place, dtype=x.dtype) + + def forward( + self, x: paddle.Tensor, offset: Union[int, paddle.Tensor] = 0 + ) -> tuple[paddle.Tensor, paddle.Tensor]: + """Add positional encoding. + + Args: + x (torch.Tensor): Input tensor (batch, time, `*`). + + Returns: + torch.Tensor: Encoded tensor (batch, time, `*`). + + """ + self.extend_pe(x) + x = x * self.xscale + pos_emb = self.position_encoding(size=x.shape[1], offset=offset) + return self.dropout(x), self.dropout(pos_emb) + + def position_encoding( + self, offset: Union[int, paddle.Tensor], size: int + ) -> paddle.Tensor: + """For getting encoding in a streaming fashion + + Attention!!!!! + we apply dropout only once at the whole utterance level in a none + streaming way, but will call this function several times with + increasing input size in a streaming scenario, so the dropout will + be applied several times. + + Args: + offset (int or torch.tensor): start offset + size (int): required size of position encoding + + Returns: + torch.Tensor: Corresponding encoding + """ + if isinstance(offset, int): + pos_emb = self.pe[ + :, + self.pe.shape[1] // 2 + - size + - offset + + 1 : self.pe.shape[1] // 2 + + size + + offset, + ] + elif isinstance(offset, paddle.Tensor): + pos_emb = self.pe[ + :, + self.pe.shape[1] // 2 + - size + - offset + + 1 : self.pe.shape[1] // 2 + + size + + offset, + ] + return pos_emb class RelPositionalEncoding(nn.Layer): """Relative positional encoding module (new implementation). diff --git a/paddlespeech/t2s/modules/transformer/encoder_layer.py b/paddlespeech/t2s/modules/transformer/encoder_layer.py index 63494b0de..b989ffd50 100644 --- a/paddlespeech/t2s/modules/transformer/encoder_layer.py +++ b/paddlespeech/t2s/modules/transformer/encoder_layer.py @@ -15,7 +15,7 @@ """Encoder self-attention layer definition.""" import paddle from paddle import nn - +from typing import Optional class EncoderLayer(nn.Layer): """Encoder layer module. @@ -111,3 +111,118 @@ class EncoderLayer(nn.Layer): x = paddle.concat([cache, x], axis=1) return x, mask + +class ConformerEncoderLayer(paddle.nn.Layer): + """Encoder layer module. + Args: + size (int): Input dimension. + self_attn (torch.nn.Module): Self-attention module instance. + `MultiHeadedAttention` or `RelPositionMultiHeadedAttention` + instance can be used as the argument. + feed_forward (torch.nn.Module): Feed-forward module instance. + `PositionwiseFeedForward` instance can be used as the argument. + feed_forward_macaron (torch.nn.Module): Additional feed-forward module + instance. + `PositionwiseFeedForward` instance can be used as the argument. + conv_module (torch.nn.Module): Convolution module instance. + `ConvlutionModule` instance can be used as the argument. + dropout_rate (float): Dropout rate. + normalize_before (bool): + True: use layer_norm before each sub-block. + False: use layer_norm after each sub-block. + """ + + def __init__( + self, + size: int, + self_attn: paddle.nn.Layer, + feed_forward: Optional[paddle.nn.Layer] = None, + feed_forward_macaron: Optional[paddle.nn.Layer] = None, + conv_module: Optional[paddle.nn.Layer] = None, + dropout_rate: float = 0.1, + normalize_before: bool = True, + ): + """Construct an EncoderLayer object.""" + super().__init__() + self.self_attn = self_attn + self.feed_forward = feed_forward + self.feed_forward_macaron = feed_forward_macaron + self.conv_module = conv_module + self.norm_ff = paddle.nn.LayerNorm(normalized_shape=size, epsilon=1e-12) + self.norm_mha = paddle.nn.LayerNorm(normalized_shape=size, epsilon=1e-12) + if feed_forward_macaron is not None: + self.norm_ff_macaron = paddle.nn.LayerNorm( + normalized_shape=size, epsilon=1e-12 + ) + self.ff_scale = 0.5 + else: + self.ff_scale = 1.0 + if self.conv_module is not None: + self.norm_conv = paddle.nn.LayerNorm(normalized_shape=size, epsilon=1e-12) + self.norm_final = paddle.nn.LayerNorm(normalized_shape=size, epsilon=1e-12) + self.dropout = paddle.nn.Dropout(p=dropout_rate) + self.size = size + self.normalize_before = normalize_before + + def forward( + self, + x: paddle.Tensor, + mask: paddle.Tensor, + pos_emb: paddle.Tensor, + mask_pad: paddle.Tensor = paddle.ones((0, 0, 0), dtype=paddle.bool), + att_cache: paddle.Tensor = paddle.zeros((0, 0, 0, 0)), + cnn_cache: paddle.Tensor = paddle.zeros((0, 0, 0, 0)), + ) -> tuple[paddle.Tensor, paddle.Tensor, paddle.Tensor, paddle.Tensor]: + """Compute encoded features. + + Args: + x (torch.Tensor): (#batch, time, size) + mask (torch.Tensor): Mask tensor for the input (#batch, time,time), + (0, 0, 0) means fake mask. + pos_emb (torch.Tensor): positional encoding, must not be None + for ConformerEncoderLayer. + mask_pad (torch.Tensor): batch padding mask used for conv module. + (#batch, 1,time), (0, 0, 0) means fake mask. + att_cache (torch.Tensor): Cache tensor of the KEY & VALUE + (#batch=1, head, cache_t1, d_k * 2), head * d_k == size. + cnn_cache (torch.Tensor): Convolution cache in conformer layer + (#batch=1, size, cache_t2) + Returns: + torch.Tensor: Output tensor (#batch, time, size). + torch.Tensor: Mask tensor (#batch, time, time). + torch.Tensor: att_cache tensor, + (#batch=1, head, cache_t1 + time, d_k * 2). + torch.Tensor: cnn_cahce tensor (#batch, size, cache_t2). + """ + if self.feed_forward_macaron is not None: + residual = x + if self.normalize_before: + x = self.norm_ff_macaron(x) + x = residual + self.ff_scale * self.dropout(self.feed_forward_macaron(x)) + if not self.normalize_before: + x = self.norm_ff_macaron(x) + residual = x + if self.normalize_before: + x = self.norm_mha(x) + x_att, new_att_cache = self.self_attn(x, x, x, pos_emb, mask,att_cache) + x = residual + self.dropout(x_att) + if not self.normalize_before: + x = self.norm_mha(x) + new_cnn_cache = paddle.zeros([0, 0, 0], dtype=x.dtype) + if self.conv_module is not None: + residual = x + if self.normalize_before: + x = self.norm_conv(x) + x, new_cnn_cache = self.conv_module(x, mask_pad, cnn_cache) + x = residual + self.dropout(x) + if not self.normalize_before: + x = self.norm_conv(x) + residual = x + if self.normalize_before: + x = self.norm_ff(x) + x = residual + self.ff_scale * self.dropout(self.feed_forward(x)) + if not self.normalize_before: + x = self.norm_ff(x) + if self.conv_module is not None: + x = self.norm_final(x) + return x, mask, new_att_cache, new_cnn_cache \ No newline at end of file diff --git a/paddlespeech/t2s/modules/transformer/espnet.py b/paddlespeech/t2s/modules/transformer/espnet.py new file mode 100644 index 000000000..229b1bd60 --- /dev/null +++ b/paddlespeech/t2s/modules/transformer/espnet.py @@ -0,0 +1,275 @@ +import paddle + +"""Positonal Encoding Module.""" +import math +from typing import Tuple, Union + +import numpy as np + + +class PositionalEncoding(paddle.nn.Layer): + """Positional encoding. + + :param int d_model: embedding dim + :param float dropout_rate: dropout rate + :param int max_len: maximum input length + + PE(pos, 2i) = sin(pos/(10000^(2i/dmodel))) + PE(pos, 2i+1) = cos(pos/(10000^(2i/dmodel))) + """ + + def __init__( + self, + d_model: int, + dropout_rate: float, + max_len: int = 5000, + reverse: bool = False, + ): + """Construct an PositionalEncoding object.""" + super().__init__() + self.d_model = d_model + self.xscale = math.sqrt(self.d_model) + self.dropout = paddle.nn.Dropout(p=dropout_rate) + self.max_len = max_len + self.pe = paddle.zeros(self.max_len, self.d_model) + position = paddle.arange(0, self.max_len, dtype=paddle.float32).unsqueeze(1) + div_term = paddle.exp( + x=paddle.arange(0, self.d_model, 2, dtype=paddle.float32) + * -(math.log(10000.0) / self.d_model) + ) + self.pe[:, 0::2] = paddle.sin(position * div_term) + self.pe[:, 1::2] = paddle.cos(position * div_term) + self.pe = self.pe.unsqueeze(0) + + def forward( + self, x: paddle.Tensor, offset: Union[int, paddle.Tensor] = 0 + ) -> Tuple[paddle.Tensor, paddle.Tensor]: + """Add positional encoding. + + Args: + x (torch.Tensor): Input. Its shape is (batch, time, ...) + offset (int, torch.tensor): position offset + + Returns: + torch.Tensor: Encoded tensor. Its shape is (batch, time, ...) + torch.Tensor: for compatibility to RelPositionalEncoding + """ + self.pe = self.pe.to(x.place) + pos_emb = self.position_encoding(offset, x.size(1), False) + x = x * self.xscale + pos_emb + return self.dropout(x), self.dropout(pos_emb) + + def position_encoding( + self, offset: Union[int, paddle.Tensor], size: int, apply_dropout: bool = True + ) -> paddle.Tensor: + """For getting encoding in a streaming fashion + + Attention!!!!! + we apply dropout only once at the whole utterance level in a none + streaming way, but will call this function several times with + increasing input size in a streaming scenario, so the dropout will + be applied several times. + + Args: + offset (int or torch.tensor): start offset + size (int): required size of position encoding + + Returns: + torch.Tensor: Corresponding encoding + """ + if isinstance(offset, int): + assert offset + size <= self.max_len + pos_emb = self.pe[:, offset : offset + size] + elif isinstance(offset, paddle.Tensor) and offset.dim() == 0: + assert offset + size <= self.max_len + pos_emb = self.pe[:, offset : offset + size] + else: + assert paddle.compat.max(offset) + size <= self.max_len + index = offset.unsqueeze(1) + paddle.arange(0, size).to(offset.place) + flag = index > 0 + index = index * flag + pos_emb = paddle.nn.functional.embedding(index, self.pe[0]) + if apply_dropout: + pos_emb = self.dropout(pos_emb) + return pos_emb + + +class RelPositionalEncoding(PositionalEncoding): + """Relative positional encoding module. + See : Appendix B in https://arxiv.org/abs/1901.02860 + Args: + d_model (int): Embedding dimension. + dropout_rate (float): Dropout rate. + max_len (int): Maximum input length. + """ + + def __init__(self, d_model: int, dropout_rate: float, max_len: int = 5000): + """Initialize class.""" + super().__init__(d_model, dropout_rate, max_len, reverse=True) + + def forward( + self, x: paddle.Tensor, offset: Union[int, paddle.Tensor] = 0 + ) -> Tuple[paddle.Tensor, paddle.Tensor]: + """Compute positional encoding. + Args: + x (torch.Tensor): Input tensor (batch, time, `*`). + Returns: + torch.Tensor: Encoded tensor (batch, time, `*`). + torch.Tensor: Positional embedding tensor (1, time, `*`). + """ + self.pe = self.pe.to(x.place) + x = x * self.xscale + pos_emb = self.position_encoding(offset, x.size(1), False) + return self.dropout(x), self.dropout(pos_emb) + + +class WhisperPositionalEncoding(PositionalEncoding): + """Sinusoids position encoding used in openai-whisper.encoder""" + + def __init__(self, d_model: int, dropout_rate: float, max_len: int = 1500): + super().__init__(d_model, dropout_rate, max_len) + self.xscale = 1.0 + log_timescale_increment = np.log(10000) / (d_model // 2 - 1) + inv_timescales = paddle.exp( + x=-log_timescale_increment * paddle.arange(d_model // 2) + ) + scaled_time = ( + paddle.arange(max_len)[:, np.newaxis] * inv_timescales[np.newaxis, :] + ) + pe = paddle.cat([paddle.sin(scaled_time), paddle.cos(scaled_time)], dim=1) + delattr(self, "pe") + self.register_buffer(name="pe", tensor=pe.unsqueeze(0)) + + +class LearnablePositionalEncoding(PositionalEncoding): + """Learnable position encoding used in openai-whisper.decoder""" + + def __init__(self, d_model: int, dropout_rate: float, max_len: int = 448): + super().__init__(d_model, dropout_rate, max_len) + self.pe = paddle.nn.parameter.Parameter(paddle.empty(1, max_len, d_model)) + self.xscale = 1.0 + + +class NoPositionalEncoding(paddle.nn.Layer): + """No position encoding""" + + def __init__(self, d_model: int, dropout_rate: float): + super().__init__() + self.d_model = d_model + self.dropout = paddle.nn.Dropout(p=dropout_rate) + + def forward( + self, x: paddle.Tensor, offset: Union[int, paddle.Tensor] = 0 + ) -> Tuple[paddle.Tensor, paddle.Tensor]: + """Just return zero vector for interface compatibility""" + pos_emb = paddle.zeros(1, x.size(1), self.d_model).to(x.place) + return self.dropout(x), pos_emb + + def position_encoding( + self, offset: Union[int, paddle.Tensor], size: int + ) -> paddle.Tensor: + return paddle.zeros(1, size, self.d_model) + + +class EspnetRelPositionalEncoding(paddle.nn.Layer): + """Relative positional encoding module (new implementation). + + Details can be found in https://github.com/espnet/espnet/pull/2816. + + See : Appendix B in https://arxiv.org/abs/1901.02860 + + Args: + d_model (int): Embedding dimension. + dropout_rate (float): Dropout rate. + max_len (int): Maximum input length. + + """ + + def __init__(self, d_model: int, dropout_rate: float, max_len: int = 5000): + """Construct an PositionalEncoding object.""" + super(EspnetRelPositionalEncoding, self).__init__() + self.d_model = d_model + self.xscale = math.sqrt(self.d_model) + self.dropout = paddle.nn.Dropout(p=dropout_rate) + self.pe = None + self.extend_pe(paddle.tensor(0.0).expand(1, max_len)) + + def extend_pe(self, x: paddle.Tensor): + """Reset the positional encodings.""" + if self.pe is not None: + if self.pe.size(1) >= x.size(1) * 2 - 1: + if self.pe.dtype != x.dtype or self.pe.place != x.place: + self.pe = self.pe.to(dtype=x.dtype, device=x.place) + return + pe_positive = paddle.zeros(x.size(1), self.d_model) + pe_negative = paddle.zeros(x.size(1), self.d_model) + position = paddle.arange(0, x.size(1), dtype=paddle.float32).unsqueeze(1) + div_term = paddle.exp( + x=paddle.arange(0, self.d_model, 2, dtype=paddle.float32) + * -(math.log(10000.0) / self.d_model) + ) + pe_positive[:, 0::2] = paddle.sin(position * div_term) + pe_positive[:, 1::2] = paddle.cos(position * div_term) + pe_negative[:, 0::2] = paddle.sin(-1 * position * div_term) + pe_negative[:, 1::2] = paddle.cos(-1 * position * div_term) + pe_positive = paddle.flip(x=pe_positive, axis=[0]).unsqueeze(0) + pe_negative = pe_negative[1:].unsqueeze(0) + pe = paddle.cat([pe_positive, pe_negative], dim=1) + self.pe = pe.to(device=x.place, dtype=x.dtype) + + def forward( + self, x: paddle.Tensor, offset: Union[int, paddle.Tensor] = 0 + ) -> Tuple[paddle.Tensor, paddle.Tensor]: + """Add positional encoding. + + Args: + x (torch.Tensor): Input tensor (batch, time, `*`). + + Returns: + torch.Tensor: Encoded tensor (batch, time, `*`). + + """ + self.extend_pe(x) + x = x * self.xscale + pos_emb = self.position_encoding(size=x.size(1), offset=offset) + return self.dropout(x), self.dropout(pos_emb) + + def position_encoding( + self, offset: Union[int, paddle.Tensor], size: int + ) -> paddle.Tensor: + """For getting encoding in a streaming fashion + + Attention!!!!! + we apply dropout only once at the whole utterance level in a none + streaming way, but will call this function several times with + increasing input size in a streaming scenario, so the dropout will + be applied several times. + + Args: + offset (int or torch.tensor): start offset + size (int): required size of position encoding + + Returns: + torch.Tensor: Corresponding encoding + """ + if isinstance(offset, int): + pos_emb = self.pe[ + :, + self.pe.size(1) // 2 + - size + - offset + + 1 : self.pe.size(1) // 2 + + size + + offset, + ] + elif isinstance(offset, paddle.Tensor): + pos_emb = self.pe[ + :, + self.pe.size(1) // 2 + - size + - offset + + 1 : self.pe.size(1) // 2 + + size + + offset, + ] + return pos_emb diff --git a/paddlespeech/t2s/modules/transformer/mask.py b/paddlespeech/t2s/modules/transformer/mask.py index 71dd37975..b0207b9fe 100644 --- a/paddlespeech/t2s/modules/transformer/mask.py +++ b/paddlespeech/t2s/modules/transformer/mask.py @@ -50,3 +50,106 @@ def target_mask(ys_in_pad, ignore_id, dtype=paddle.bool): ys_mask = ys_in_pad != ignore_id m = subsequent_mask(ys_mask.shape[-1]).unsqueeze(0) return ys_mask.unsqueeze(-2) & m + +def make_pad_mask(lengths: paddle.Tensor, max_len: int = 0) -> paddle.Tensor: + """Make mask tensor containing indices of padded part. + + See description of make_non_pad_mask. + + Args: + lengths (torch.Tensor): Batch of lengths (B,). + Returns: + torch.Tensor: Mask tensor containing indices of padded part. + + Examples: + >>> lengths = [5, 3, 2] + >>> make_pad_mask(lengths) + masks = [[0, 0, 0, 0 ,0], + [0, 0, 0, 1, 1], + [0, 0, 1, 1, 1]] + """ + batch_size = lengths.shape[0] + max_len = max_len if max_len > 0 else lengths._max().item() + seq_range = paddle.arange(0, max_len, dtype=paddle.int32) + seq_range_expand = seq_range.unsqueeze(0).expand([batch_size, max_len]) + seq_length_expand = lengths.unsqueeze(-1) + mask = seq_range_expand >= seq_length_expand + return mask + +def add_optional_chunk_mask( + xs: paddle.Tensor, + masks: paddle.Tensor, + use_dynamic_chunk: bool, + use_dynamic_left_chunk: bool, + decoding_chunk_size: int, + static_chunk_size: int, + num_decoding_left_chunks: int, + enable_full_context: bool = True, +): + """Apply optional mask for encoder. + + Args: + xs (torch.Tensor): padded input, (B, L, D), L for max length + mask (torch.Tensor): mask for xs, (B, 1, L) + use_dynamic_chunk (bool): whether to use dynamic chunk or not + use_dynamic_left_chunk (bool): whether to use dynamic left chunk for + training. + decoding_chunk_size (int): decoding chunk size for dynamic chunk, it's + 0: default for training, use random dynamic chunk. + <0: for decoding, use full chunk. + >0: for decoding, use fixed chunk size as set. + static_chunk_size (int): chunk size for static chunk training/decoding + if it's greater than 0, if use_dynamic_chunk is true, + this parameter will be ignored + num_decoding_left_chunks: number of left chunks, this is for decoding, + the chunk size is decoding_chunk_size. + >=0: use num_decoding_left_chunks + <0: use all left chunks + enable_full_context (bool): + True: chunk size is either [1, 25] or full context(max_len) + False: chunk size ~ U[1, 25] + + Returns: + torch.Tensor: chunk mask of the input xs. + """ + if use_dynamic_chunk: + max_len = xs.size(1) + if decoding_chunk_size < 0: + chunk_size = max_len + num_left_chunks = -1 + elif decoding_chunk_size > 0: + chunk_size = decoding_chunk_size + num_left_chunks = num_decoding_left_chunks + else: + chunk_size = paddle.randint(low=1, high=max_len, shape=(1,)).item() + num_left_chunks = -1 + if chunk_size > max_len // 2 and enable_full_context: + chunk_size = max_len + else: + chunk_size = chunk_size % 25 + 1 + if use_dynamic_left_chunk: + max_left_chunks = (max_len - 1) // chunk_size + num_left_chunks = paddle.randint( + low=0, high=max_left_chunks, shape=(1,) + ).item() + chunk_masks = subsequent_chunk_mask( + xs.size(1), chunk_size, num_left_chunks, xs.place + ) + chunk_masks = chunk_masks.unsqueeze(0) + chunk_masks = masks & chunk_masks + elif static_chunk_size > 0: + num_left_chunks = num_decoding_left_chunks + chunk_masks = subsequent_chunk_mask( + xs.size(1), static_chunk_size, num_left_chunks, xs.place + ) + chunk_masks = chunk_masks.unsqueeze(0) + chunk_masks = masks & chunk_masks + else: + chunk_masks = masks + assert chunk_masks.dtype == paddle.bool + if (chunk_masks.sum(axis=-1) == 0).sum().item() != 0: + print( + "get chunk_masks all false at some timestep, force set to true, make sure they are masked in futuer computation!" + ) + chunk_masks[chunk_masks.sum(axis=-1) == 0] = True + return chunk_masks \ No newline at end of file diff --git a/paddlespeech/t2s/modules/transformer/subsampling.py b/paddlespeech/t2s/modules/transformer/subsampling.py index a17278c0b..ca46569bb 100644 --- a/paddlespeech/t2s/modules/transformer/subsampling.py +++ b/paddlespeech/t2s/modules/transformer/subsampling.py @@ -15,10 +15,66 @@ """Subsampling layer definition.""" import paddle from paddle import nn - +from typing import Union from paddlespeech.t2s.modules.transformer.embedding import PositionalEncoding +class BaseSubsampling(paddle.nn.Layer): + def __init__(self): + super().__init__() + self.right_context = 0 + self.subsampling_rate = 1 + + def position_encoding( + self, offset: Union[int, paddle.Tensor], size: int + ) -> paddle.Tensor: + return self.pos_enc.position_encoding(offset, size) +class LinearNoSubsampling(BaseSubsampling): + """Linear transform the input without subsampling + + Args: + idim (int): Input dimension. + odim (int): Output dimension. + dropout_rate (float): Dropout rate. + + """ + + def __init__( + self, idim: int, odim: int, dropout_rate: float, pos_enc_class: paddle.nn.Layer + ): + """Construct an linear object.""" + super().__init__() + self.out = paddle.nn.Sequential( + paddle.nn.Linear(in_features=idim, out_features=odim), + paddle.nn.LayerNorm(normalized_shape=odim, epsilon=1e-05), + paddle.nn.Dropout(p=dropout_rate), + ) + self.pos_enc = pos_enc_class + self.right_context = 0 + self.subsampling_rate = 1 + + def forward( + self, + x: paddle.Tensor, + x_mask: paddle.Tensor, + offset: Union[int, paddle.Tensor] = 0, + ) -> tuple[paddle.Tensor, paddle.Tensor, paddle.Tensor]: + """Input x. + + Args: + x (paddle.Tensor): Input tensor (#batch, time, idim). + x_mask (torch.Tensor): Input mask (#batch, 1, time). + + Returns: + paddle.Tensor: linear input tensor (#batch, time', odim), + where time' = time . + paddle.Tensor: linear input mask (#batch, 1, time'), + where time' = time . + + """ + x = self.out(x) + x, pos_emb = self.pos_enc(x, offset) + return x, pos_emb, x_mask class Conv2dSubsampling(nn.Layer): """Convolutional 2D subsampling (to 1/4 length). diff --git a/paddlespeech/t2s/modules/transformer/upsample_encoder.py b/paddlespeech/t2s/modules/transformer/upsample_encoder.py new file mode 100644 index 000000000..dfdcb1620 --- /dev/null +++ b/paddlespeech/t2s/modules/transformer/upsample_encoder.py @@ -0,0 +1,343 @@ +import paddle + +"""Encoder definition.""" +from typing import Tuple +import paddle.nn.functional as F +from paddlespeech.t2s.modules.transformer.convolution import ConvolutionModule +from paddlespeech.t2s.modules.transformer.encoder_layer import ConformerEncoderLayer +from paddlespeech.t2s.modules.transformer.positionwise_feed_forward import PositionwiseFeedForward +from paddlespeech.t2s.models.CosyVoice.class_utils import (COSYVOICE_ACTIVATION_CLASSES, + COSYVOICE_ATTENTION_CLASSES, + COSYVOICE_EMB_CLASSES, + COSYVOICE_SUBSAMPLE_CLASSES) +from paddlespeech.t2s.modules.transformer.mask import add_optional_chunk_mask, make_pad_mask + + +class Upsample1D(paddle.nn.Layer): + """A 1D upsampling layer with an optional convolution. + + Parameters: + channels (`int`): + number of channels in the inputs and outputs. + use_conv (`bool`, default `False`): + option to use a convolution. + use_conv_transpose (`bool`, default `False`): + option to use a convolution transpose. + out_channels (`int`, optional): + number of output channels. Defaults to `channels`. + """ + + def __init__(self, channels: int, out_channels: int, stride: int = 2): + super().__init__() + self.channels = channels + self.out_channels = out_channels + self.stride = stride + self.conv = paddle.nn.Conv1D( + self.channels, self.out_channels, stride * 2 + 1, stride=1, padding=0 + ) + + def forward( + self, inputs: paddle.Tensor, input_lengths: paddle.Tensor + ) -> Tuple[paddle.Tensor, paddle.Tensor]: + inputs = inputs.unsqueeze(2) + outputs = paddle.nn.functional.interpolate( + x=inputs, scale_factor=[1, float(self.stride)],mode="nearest" + ) + outputs = outputs.squeeze(2) + outputs = F.pad(outputs, [self.stride * 2, 0], value=0.0) + outputs = self.conv(outputs) + return outputs, input_lengths * self.stride + + +class PreLookaheadLayer(paddle.nn.Layer): + def __init__(self, channels: int, pre_lookahead_len: int = 1): + super().__init__() + self.channels = channels + self.pre_lookahead_len = pre_lookahead_len + self.conv1 = paddle.nn.Conv1D( + channels, channels, kernel_size=pre_lookahead_len + 1, stride=1, padding=0 + ) + self.conv2 = paddle.nn.Conv1D( + channels, channels, kernel_size=3, stride=1, padding=0 + ) + + def forward( + self, inputs: paddle.Tensor, context: paddle.Tensor = paddle.zeros([0, 0, 0]) + ) -> paddle.Tensor: + """ + inputs: (batch_size, seq_len, channels) + """ + outputs = paddle.transpose(inputs, perm=[0, 2, 1]).contiguous() + context = paddle.transpose(context, perm=[0, 2, 1]).contiguous() + if context.shape[2] == 0: + outputs = F.pad( + outputs, [0, self.pre_lookahead_len], mode="constant", value=0.0 + ) + else: + assert ( + self.training is False + ), "you have passed context, make sure that you are running inference mode" + assert context.shape[2] == self.pre_lookahead_len + outputs = F.pad( + paddle.cat([outputs, context], dim=2), + [0, self.pre_lookahead_len - context.shape[2]], + mode="constant", + value=0.0, + ) + outputs = paddle.nn.functional.leaky_relu(x=self.conv1(outputs)) + outputs = F.pad( + outputs, [self.conv2._kernel_size[0] - 1, 0], mode="constant", value=0.0 + ) + outputs = self.conv2(outputs) + outputs = paddle.transpose(outputs, perm=[0, 2, 1]).contiguous() + + outputs = outputs + inputs + return outputs + + +class UpsampleConformerEncoder(paddle.nn.Layer): + def __init__( + self, + input_size: int, + output_size: int = 256, + attention_heads: int = 4, + linear_units: int = 2048, + num_blocks: int = 6, + dropout_rate: float = 0.1, + positional_dropout_rate: float = 0.1, + attention_dropout_rate: float = 0.0, + input_layer: str = "conv2d", + pos_enc_layer_type: str = "rel_pos", + normalize_before: bool = True, + static_chunk_size: int = 0, + use_dynamic_chunk: bool = False, + global_cmvn: paddle.nn.Layer = None, + use_dynamic_left_chunk: bool = False, + positionwise_conv_kernel_size: int = 1, + macaron_style: bool = True, + selfattention_layer_type: str = "rel_selfattn", + activation_type: str = "swish", + use_cnn_module: bool = True, + cnn_module_kernel: int = 15, + causal: bool = False, + cnn_module_norm: str = "batch_norm", + key_bias: bool = True, + gradient_checkpointing: bool = False, + ): + """ + Args: + input_size (int): input dim + output_size (int): dimension of attention + attention_heads (int): the number of heads of multi head attention + linear_units (int): the hidden units number of position-wise feed + forward + num_blocks (int): the number of decoder blocks + dropout_rate (float): dropout rate + attention_dropout_rate (float): dropout rate in attention + positional_dropout_rate (float): dropout rate after adding + positional encoding + input_layer (str): input layer type. + optional [linear, conv2d, conv2d6, conv2d8] + pos_enc_layer_type (str): Encoder positional encoding layer type. + opitonal [abs_pos, scaled_abs_pos, rel_pos, no_pos] + normalize_before (bool): + True: use layer_norm before each sub-block of a layer. + False: use layer_norm after each sub-block of a layer. + static_chunk_size (int): chunk size for static chunk training and + decoding + use_dynamic_chunk (bool): whether use dynamic chunk size for + training or not, You can only use fixed chunk(chunk_size > 0) + or dyanmic chunk size(use_dynamic_chunk = True) + global_cmvn (Optional[torch.nn.Module]): Optional GlobalCMVN module + use_dynamic_left_chunk (bool): whether use dynamic left chunk in + dynamic chunk training + key_bias: whether use bias in attention.linear_k, False for whisper models. + gradient_checkpointing: rerunning a forward-pass segment for each + checkpointed segment during backward. + """ + super().__init__() + self._output_size = output_size + self.global_cmvn = global_cmvn + self.embed = COSYVOICE_SUBSAMPLE_CLASSES[input_layer]( + input_size, + output_size, + dropout_rate, + COSYVOICE_EMB_CLASSES[pos_enc_layer_type]( + output_size, positional_dropout_rate + ), + ) + self.normalize_before = normalize_before + self.after_norm = paddle.nn.LayerNorm( + normalized_shape=output_size, epsilon=1e-05 + ) + self.static_chunk_size = static_chunk_size + self.use_dynamic_chunk = use_dynamic_chunk + self.use_dynamic_left_chunk = use_dynamic_left_chunk + self.gradient_checkpointing = gradient_checkpointing + activation = COSYVOICE_ACTIVATION_CLASSES[activation_type]() + encoder_selfattn_layer_args = ( + attention_heads, + output_size, + attention_dropout_rate, + False, + ) + positionwise_layer_args = (output_size, linear_units, dropout_rate, activation) + convolution_layer_args = ( + output_size, + cnn_module_kernel, + activation, + cnn_module_norm, + causal, + ) + self.pre_lookahead_layer = PreLookaheadLayer(channels=512, pre_lookahead_len=3) + self.encoders = paddle.nn.LayerList( + sublayers=[ + ConformerEncoderLayer( + output_size, + COSYVOICE_ATTENTION_CLASSES[selfattention_layer_type]( + *encoder_selfattn_layer_args + ), + PositionwiseFeedForward(*positionwise_layer_args), + PositionwiseFeedForward(*positionwise_layer_args) + if macaron_style + else None, + ConvolutionModule(*convolution_layer_args) + if use_cnn_module + else None, + dropout_rate, + normalize_before, + ) + for _ in range(num_blocks) + ] + ) + self.up_layer = Upsample1D(channels=512, out_channels=512, stride=2) + self.up_embed = COSYVOICE_SUBSAMPLE_CLASSES[input_layer]( + input_size, + output_size, + dropout_rate, + COSYVOICE_EMB_CLASSES[pos_enc_layer_type]( + output_size, positional_dropout_rate + ), + ) + self.up_encoders = paddle.nn.LayerList( + sublayers=[ + ConformerEncoderLayer( + output_size, + COSYVOICE_ATTENTION_CLASSES[selfattention_layer_type]( + *encoder_selfattn_layer_args + ), + PositionwiseFeedForward(*positionwise_layer_args), + PositionwiseFeedForward(*positionwise_layer_args) + if macaron_style + else None, + ConvolutionModule(*convolution_layer_args) + if use_cnn_module + else None, + dropout_rate, + normalize_before, + ) + for _ in range(4) + ] + ) + + def output_size(self) -> int: + return self._output_size + + def forward( + self, + xs: paddle.Tensor, + xs_lens: paddle.Tensor, + context: paddle.Tensor = paddle.zeros([0, 0, 0]), + decoding_chunk_size: int = 0, + num_decoding_left_chunks: int = -1, + streaming: bool = False, + ) -> Tuple[paddle.Tensor, paddle.Tensor]: + """Embed positions in tensor. + + Args: + xs: padded input tensor (B, T, D) + xs_lens: input length (B) + decoding_chunk_size: decoding chunk size for dynamic chunk + 0: default for training, use random dynamic chunk. + <0: for decoding, use full chunk. + >0: for decoding, use fixed chunk size as set. + num_decoding_left_chunks: number of left chunks, this is for decoding, + the chunk size is decoding_chunk_size. + >=0: use num_decoding_left_chunks + <0: use all left chunks + Returns: + encoder output tensor xs, and subsampled masks + xs: padded output tensor (B, T' ~= T/subsample_rate, D) + masks: torch.Tensor batch padding mask after subsample + (B, 1, T' ~= T/subsample_rate) + NOTE(xcsong): + We pass the `__call__` method of the modules instead of `forward` to the + checkpointing API because `__call__` attaches all the hooks of the module. + https://discuss.pytorch.org/t/any-different-between-model-input-and-model-forward-input/3690/2 + """ + T = xs.shape[1] + masks = ~make_pad_mask(xs_lens, T).unsqueeze(1) + if self.global_cmvn is not None: + xs = self.global_cmvn(xs) + xs, pos_emb, masks = self.embed(xs, masks) + if context.shape[1] != 0: + assert ( + self.training is False + ), "you have passed context, make sure that you are running inference mode" + context_masks = paddle.ones(1, 1, context.shape[1]).to(masks) + context, _, _ = self.embed(context, context_masks, offset=xs.shape[1]) + mask_pad = masks + chunk_masks = add_optional_chunk_mask( + xs, + masks, + False, + False, + 0, + self.static_chunk_size if streaming is True else 0, + -1, + ) + xs = self.pre_lookahead_layer(xs, context=context) + xs = self.forward_layers(xs, chunk_masks, pos_emb, mask_pad) + # + xs = paddle.transpose(xs, perm=[0, 2, 1]).contiguous() + xs, xs_lens = self.up_layer(xs, xs_lens) + xs = paddle.transpose(xs, perm=[0, 2, 1]).contiguous() + T = xs.shape[1] + masks = ~make_pad_mask(xs_lens, T).unsqueeze(1) + xs, pos_emb, masks = self.up_embed(xs, masks) + mask_pad = masks + chunk_masks = add_optional_chunk_mask( + xs, + masks, + False, + False, + 0, + self.static_chunk_size * self.up_layer.stride if streaming is True else 0, + -1, + ) + xs = self.forward_up_layers(xs, chunk_masks, pos_emb, mask_pad) + if self.normalize_before: + xs = self.after_norm(xs) + return xs, masks + + def forward_layers( + self, + xs: paddle.Tensor, + chunk_masks: paddle.Tensor, + pos_emb: paddle.Tensor, + mask_pad: paddle.Tensor, + ) -> paddle.Tensor: + for layer in self.encoders: + xs, chunk_masks, _, _ = layer(xs, chunk_masks, pos_emb, mask_pad) + return xs + + def forward_up_layers( + self, + xs: paddle.Tensor, + chunk_masks: paddle.Tensor, + pos_emb: paddle.Tensor, + mask_pad: paddle.Tensor, + ) -> paddle.Tensor: + for layer in self.up_encoders: + xs, chunk_masks, _, _ = layer(xs, chunk_masks, pos_emb, mask_pad) + return xs