Team Ai
Apppublic

KUI71/ACE-Step

sourceHugging Faceapache-2.0updated 1y agoView on Hugging Face
0likes
lyric_encoder.py1071 linesDownload Raw Back to lyrics_utils
1from typing import Optional, Tuple, Union2import math3import torch4from torch import nn5 6class ConvolutionModule(nn.Module):7    """ConvolutionModule in Conformer model."""8 9    def __init__(self,10                 channels: int,11                 kernel_size: int = 15,12                 activation: nn.Module = nn.ReLU(),13                 norm: str = "batch_norm",14                 causal: bool = False,15                 bias: bool = True):16        """Construct an ConvolutionModule object.17        Args:18            channels (int): The number of channels of conv layers.19            kernel_size (int): Kernel size of conv layers.20            causal (int): Whether use causal convolution or not21        """22        super().__init__()23 24        self.pointwise_conv1 = nn.Conv1d(25            channels,26            2 * channels,27            kernel_size=1,28            stride=1,29            padding=0,30            bias=bias,31        )32        # self.lorder is used to distinguish if it's a causal convolution,33        # if self.lorder > 0: it's a causal convolution, the input will be34        #    padded with self.lorder frames on the left in forward.35        # else: it's a symmetrical convolution36        if causal:37            padding = 038            self.lorder = kernel_size - 139        else:40            # kernel_size should be an odd number for none causal convolution41            assert (kernel_size - 1) % 2 == 042            padding = (kernel_size - 1) // 243            self.lorder = 044        self.depthwise_conv = nn.Conv1d(45            channels,46            channels,47            kernel_size,48            stride=1,49            padding=padding,50            groups=channels,51            bias=bias,52        )53 54        assert norm in ['batch_norm', 'layer_norm']55        if norm == "batch_norm":56            self.use_layer_norm = False57            self.norm = nn.BatchNorm1d(channels)58        else:59            self.use_layer_norm = True60            self.norm = nn.LayerNorm(channels)61 62        self.pointwise_conv2 = nn.Conv1d(63            channels,64            channels,65            kernel_size=1,66            stride=1,67            padding=0,68            bias=bias,69        )70        self.activation = activation71 72    def forward(73        self,74        x: torch.Tensor,75        mask_pad: torch.Tensor = torch.ones((0, 0, 0), dtype=torch.bool),76        cache: torch.Tensor = torch.zeros((0, 0, 0)),77    ) -> Tuple[torch.Tensor, torch.Tensor]:78        """Compute convolution module.79        Args:80            x (torch.Tensor): Input tensor (#batch, time, channels).81            mask_pad (torch.Tensor): used for batch padding (#batch, 1, time),82                (0, 0, 0) means fake mask.83            cache (torch.Tensor): left context cache, it is only84                used in causal convolution (#batch, channels, cache_t),85                (0, 0, 0) meas fake cache.86        Returns:87            torch.Tensor: Output tensor (#batch, time, channels).88        """89        # exchange the temporal dimension and the feature dimension90        x = x.transpose(1, 2)  # (#batch, channels, time)91 92        # mask batch padding93        if mask_pad.size(2) > 0:  # time > 094            x.masked_fill_(~mask_pad, 0.0)95 96        if self.lorder > 0:97            if cache.size(2) == 0:  # cache_t == 098                x = nn.functional.pad(x, (self.lorder, 0), 'constant', 0.0)99            else:100                assert cache.size(0) == x.size(0)  # equal batch101                assert cache.size(1) == x.size(1)  # equal channel102                x = torch.cat((cache, x), dim=2)103            assert (x.size(2) > self.lorder)104            new_cache = x[:, :, -self.lorder:]105        else:106            # It's better we just return None if no cache is required,107            # However, for JIT export, here we just fake one tensor instead of108            # None.109            new_cache = torch.zeros((0, 0, 0), dtype=x.dtype, device=x.device)110 111        # GLU mechanism112        x = self.pointwise_conv1(x)  # (batch, 2*channel, dim)113        x = nn.functional.glu(x, dim=1)  # (batch, channel, dim)114 115        # 1D Depthwise Conv116        x = self.depthwise_conv(x)117        if self.use_layer_norm:118            x = x.transpose(1, 2)119        x = self.activation(self.norm(x))120        if self.use_layer_norm:121            x = x.transpose(1, 2)122        x = self.pointwise_conv2(x)123        # mask batch padding124        if mask_pad.size(2) > 0:  # time > 0125            x.masked_fill_(~mask_pad, 0.0)126 127        return x.transpose(1, 2), new_cache128 129class PositionwiseFeedForward(torch.nn.Module):130    """Positionwise feed forward layer.131 132    FeedForward are appied on each position of the sequence.133    The output dim is same with the input dim.134 135    Args:136        idim (int): Input dimenstion.137        hidden_units (int): The number of hidden units.138        dropout_rate (float): Dropout rate.139        activation (torch.nn.Module): Activation function140    """141 142    def __init__(143            self,144            idim: int,145            hidden_units: int,146            dropout_rate: float,147            activation: torch.nn.Module = torch.nn.ReLU(),148    ):149        """Construct a PositionwiseFeedForward object."""150        super(PositionwiseFeedForward, self).__init__()151        self.w_1 = torch.nn.Linear(idim, hidden_units)152        self.activation = activation153        self.dropout = torch.nn.Dropout(dropout_rate)154        self.w_2 = torch.nn.Linear(hidden_units, idim)155 156    def forward(self, xs: torch.Tensor) -> torch.Tensor:157        """Forward function.158 159        Args:160            xs: input tensor (B, L, D)161        Returns:162            output tensor, (B, L, D)163        """164        return self.w_2(self.dropout(self.activation(self.w_1(xs))))165 166class Swish(torch.nn.Module):167    """Construct an Swish object."""168 169    def forward(self, x: torch.Tensor) -> torch.Tensor:170        """Return Swish activation function."""171        return x * torch.sigmoid(x)172 173class MultiHeadedAttention(nn.Module):174    """Multi-Head Attention layer.175 176    Args:177        n_head (int): The number of heads.178        n_feat (int): The number of features.179        dropout_rate (float): Dropout rate.180 181    """182 183    def __init__(self,184                 n_head: int,185                 n_feat: int,186                 dropout_rate: float,187                 key_bias: bool = True):188        """Construct an MultiHeadedAttention object."""189        super().__init__()190        assert n_feat % n_head == 0191        # We assume d_v always equals d_k192        self.d_k = n_feat // n_head193        self.h = n_head194        self.linear_q = nn.Linear(n_feat, n_feat)195        self.linear_k = nn.Linear(n_feat, n_feat, bias=key_bias)196        self.linear_v = nn.Linear(n_feat, n_feat)197        self.linear_out = nn.Linear(n_feat, n_feat)198        self.dropout = nn.Dropout(p=dropout_rate)199 200    def forward_qkv(201        self, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor202    ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:203        """Transform query, key and value.204 205        Args:206            query (torch.Tensor): Query tensor (#batch, time1, size).207            key (torch.Tensor): Key tensor (#batch, time2, size).208            value (torch.Tensor): Value tensor (#batch, time2, size).209 210        Returns:211            torch.Tensor: Transformed query tensor, size212                (#batch, n_head, time1, d_k).213            torch.Tensor: Transformed key tensor, size214                (#batch, n_head, time2, d_k).215            torch.Tensor: Transformed value tensor, size216                (#batch, n_head, time2, d_k).217 218        """219        n_batch = query.size(0)220        q = self.linear_q(query).view(n_batch, -1, self.h, self.d_k)221        k = self.linear_k(key).view(n_batch, -1, self.h, self.d_k)222        v = self.linear_v(value).view(n_batch, -1, self.h, self.d_k)223        q = q.transpose(1, 2)  # (batch, head, time1, d_k)224        k = k.transpose(1, 2)  # (batch, head, time2, d_k)225        v = v.transpose(1, 2)  # (batch, head, time2, d_k)226        return q, k, v227 228    def forward_attention(229        self,230        value: torch.Tensor,231        scores: torch.Tensor,232        mask: torch.Tensor = torch.ones((0, 0, 0), dtype=torch.bool)233    ) -> torch.Tensor:234        """Compute attention context vector.235 236        Args:237            value (torch.Tensor): Transformed value, size238                (#batch, n_head, time2, d_k).239            scores (torch.Tensor): Attention score, size240                (#batch, n_head, time1, time2).241            mask (torch.Tensor): Mask, size (#batch, 1, time2) or242                (#batch, time1, time2), (0, 0, 0) means fake mask.243 244        Returns:245            torch.Tensor: Transformed value (#batch, time1, d_model)246                weighted by the attention score (#batch, time1, time2).247 248        """249        n_batch = value.size(0)250    251        if mask.size(2) > 0:  # time2 > 0252            mask = mask.unsqueeze(1).eq(0)  # (batch, 1, *, time2)253            # For last chunk, time2 might be larger than scores.size(-1)254            mask = mask[:, :, :, :scores.size(-1)]  # (batch, 1, *, time2)255            scores = scores.masked_fill(mask, -float('inf'))256            attn = torch.softmax(scores, dim=-1).masked_fill(257                mask, 0.0)  # (batch, head, time1, time2)258 259        else:260            attn = torch.softmax(scores, dim=-1)  # (batch, head, time1, time2)261 262        p_attn = self.dropout(attn)263        x = torch.matmul(p_attn, value)  # (batch, head, time1, d_k)264        x = (x.transpose(1, 2).contiguous().view(n_batch, -1,265                                                 self.h * self.d_k)266             )  # (batch, time1, d_model)267 268        return self.linear_out(x)  # (batch, time1, d_model)269 270    def forward(271        self,272        query: torch.Tensor,273        key: torch.Tensor,274        value: torch.Tensor,275        mask: torch.Tensor = torch.ones((0, 0, 0), dtype=torch.bool),276        pos_emb: torch.Tensor = torch.empty(0),277        cache: torch.Tensor = torch.zeros((0, 0, 0, 0))278    ) -> Tuple[torch.Tensor, torch.Tensor]:279        """Compute scaled dot product attention.280 281        Args:282            query (torch.Tensor): Query tensor (#batch, time1, size).283            key (torch.Tensor): Key tensor (#batch, time2, size).284            value (torch.Tensor): Value tensor (#batch, time2, size).285            mask (torch.Tensor): Mask tensor (#batch, 1, time2) or286                (#batch, time1, time2).287                1.When applying cross attention between decoder and encoder,288                the batch padding mask for input is in (#batch, 1, T) shape.289                2.When applying self attention of encoder,290                the mask is in (#batch, T, T)  shape.291                3.When applying self attention of decoder,292                the mask is in (#batch, L, L)  shape.293                4.If the different position in decoder see different block294                of the encoder, such as Mocha, the passed in mask could be295                in (#batch, L, T) shape. But there is no such case in current296                CosyVoice.297            cache (torch.Tensor): Cache tensor (1, head, cache_t, d_k * 2),298                where `cache_t == chunk_size * num_decoding_left_chunks`299                and `head * d_k == size`300 301 302        Returns:303            torch.Tensor: Output tensor (#batch, time1, d_model).304            torch.Tensor: Cache tensor (1, head, cache_t + time1, d_k * 2)305                where `cache_t == chunk_size * num_decoding_left_chunks`306                and `head * d_k == size`307 308        """309        q, k, v = self.forward_qkv(query, key, value)310        if cache.size(0) > 0:311            key_cache, value_cache = torch.split(cache,312                                                 cache.size(-1) // 2,313                                                 dim=-1)314            k = torch.cat([key_cache, k], dim=2)315            v = torch.cat([value_cache, v], dim=2)316        new_cache = torch.cat((k, v), dim=-1)317 318        scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k)319        return self.forward_attention(v, scores, mask), new_cache320 321 322class RelPositionMultiHeadedAttention(MultiHeadedAttention):323    """Multi-Head Attention layer with relative position encoding.324    Paper: https://arxiv.org/abs/1901.02860325    Args:326        n_head (int): The number of heads.327        n_feat (int): The number of features.328        dropout_rate (float): Dropout rate.329    """330 331    def __init__(self,332                 n_head: int,333                 n_feat: int,334                 dropout_rate: float,335                 key_bias: bool = True):336        """Construct an RelPositionMultiHeadedAttention object."""337        super().__init__(n_head, n_feat, dropout_rate, key_bias)338        # linear transformation for positional encoding339        self.linear_pos = nn.Linear(n_feat, n_feat, bias=False)340        # these two learnable bias are used in matrix c and matrix d341        # as described in https://arxiv.org/abs/1901.02860 Section 3.3342        self.pos_bias_u = nn.Parameter(torch.Tensor(self.h, self.d_k))343        self.pos_bias_v = nn.Parameter(torch.Tensor(self.h, self.d_k))344        torch.nn.init.xavier_uniform_(self.pos_bias_u)345        torch.nn.init.xavier_uniform_(self.pos_bias_v)346 347    def rel_shift(self, x: torch.Tensor) -> torch.Tensor:348        """Compute relative positional encoding.349 350        Args:351            x (torch.Tensor): Input tensor (batch, head, time1, 2*time1-1).352            time1 means the length of query vector.353 354        Returns:355            torch.Tensor: Output tensor.356 357        """358        zero_pad = torch.zeros((x.size()[0], x.size()[1], x.size()[2], 1),359                               device=x.device,360                               dtype=x.dtype)361        x_padded = torch.cat([zero_pad, x], dim=-1)362 363        x_padded = x_padded.view(x.size()[0],364                                 x.size()[1],365                                 x.size(3) + 1, x.size(2))366        x = x_padded[:, :, 1:].view_as(x)[367            :, :, :, : x.size(-1) // 2 + 1368        ]  # only keep the positions from 0 to time2369        return x370 371    def forward(372        self,373        query: torch.Tensor,374        key: torch.Tensor,375        value: torch.Tensor,376        mask: torch.Tensor = torch.ones((0, 0, 0), dtype=torch.bool),377        pos_emb: torch.Tensor = torch.empty(0),378        cache: torch.Tensor = torch.zeros((0, 0, 0, 0))379    ) -> Tuple[torch.Tensor, torch.Tensor]:380        """Compute 'Scaled Dot Product Attention' with rel. positional encoding.381        Args:382            query (torch.Tensor): Query tensor (#batch, time1, size).383            key (torch.Tensor): Key tensor (#batch, time2, size).384            value (torch.Tensor): Value tensor (#batch, time2, size).385            mask (torch.Tensor): Mask tensor (#batch, 1, time2) or386                (#batch, time1, time2), (0, 0, 0) means fake mask.387            pos_emb (torch.Tensor): Positional embedding tensor388                (#batch, time2, size).389            cache (torch.Tensor): Cache tensor (1, head, cache_t, d_k * 2),390                where `cache_t == chunk_size * num_decoding_left_chunks`391                and `head * d_k == size`392        Returns:393            torch.Tensor: Output tensor (#batch, time1, d_model).394            torch.Tensor: Cache tensor (1, head, cache_t + time1, d_k * 2)395                where `cache_t == chunk_size * num_decoding_left_chunks`396                and `head * d_k == size`397        """398        q, k, v = self.forward_qkv(query, key, value)399        q = q.transpose(1, 2)  # (batch, time1, head, d_k)400 401        if cache.size(0) > 0:402            key_cache, value_cache = torch.split(cache,403                                                 cache.size(-1) // 2,404                                                 dim=-1)405            k = torch.cat([key_cache, k], dim=2)406            v = torch.cat([value_cache, v], dim=2)407        # NOTE(xcsong): We do cache slicing in encoder.forward_chunk, since it's408        #   non-trivial to calculate `next_cache_start` here.409        new_cache = torch.cat((k, v), dim=-1)410 411        n_batch_pos = pos_emb.size(0)412        p = self.linear_pos(pos_emb).view(n_batch_pos, -1, self.h, self.d_k)413        p = p.transpose(1, 2)  # (batch, head, time1, d_k)414 415        # (batch, head, time1, d_k)416        q_with_bias_u = (q + self.pos_bias_u).transpose(1, 2)417        # (batch, head, time1, d_k)418        q_with_bias_v = (q + self.pos_bias_v).transpose(1, 2)419 420        # compute attention score421        # first compute matrix a and matrix c422        # as described in https://arxiv.org/abs/1901.02860 Section 3.3423        # (batch, head, time1, time2)424        matrix_ac = torch.matmul(q_with_bias_u, k.transpose(-2, -1))425 426        # compute matrix b and matrix d427        # (batch, head, time1, time2)428        matrix_bd = torch.matmul(q_with_bias_v, p.transpose(-2, -1))429        # NOTE(Xiang Lyu): Keep rel_shift since espnet rel_pos_emb is used430        if matrix_ac.shape != matrix_bd.shape:431            matrix_bd = self.rel_shift(matrix_bd)432 433        scores = (matrix_ac + matrix_bd) / math.sqrt(434            self.d_k)  # (batch, head, time1, time2)435 436        return self.forward_attention(v, scores, mask), new_cache437 438 439 440def subsequent_mask(441        size: int,442        device: torch.device = torch.device("cpu"),443) -> torch.Tensor:444    """Create mask for subsequent steps (size, size).445 446    This mask is used only in decoder which works in an auto-regressive mode.447    This means the current step could only do attention with its left steps.448 449    In encoder, fully attention is used when streaming is not necessary and450    the sequence is not long. In this  case, no attention mask is needed.451 452    When streaming is need, chunk-based attention is used in encoder. See453    subsequent_chunk_mask for the chunk-based attention mask.454 455    Args:456        size (int): size of mask457        str device (str): "cpu" or "cuda" or torch.Tensor.device458        dtype (torch.device): result dtype459 460    Returns:461        torch.Tensor: mask462 463    Examples:464        >>> subsequent_mask(3)465        [[1, 0, 0],466         [1, 1, 0],467         [1, 1, 1]]468    """469    arange = torch.arange(size, device=device)470    mask = arange.expand(size, size)471    arange = arange.unsqueeze(-1)472    mask = mask <= arange473    return mask474 475 476def subsequent_chunk_mask(477        size: int,478        chunk_size: int,479        num_left_chunks: int = -1,480        device: torch.device = torch.device("cpu"),481    ) -> torch.Tensor:482    """Create mask for subsequent steps (size, size) with chunk size,483       this is for streaming encoder484 485    Args:486        size (int): size of mask487        chunk_size (int): size of chunk488        num_left_chunks (int): number of left chunks489            <0: use full chunk490            >=0: use num_left_chunks491        device (torch.device): "cpu" or "cuda" or torch.Tensor.device492 493    Returns:494        torch.Tensor: mask495 496    Examples:497        >>> subsequent_chunk_mask(4, 2)498        [[1, 1, 0, 0],499         [1, 1, 0, 0],500         [1, 1, 1, 1],501         [1, 1, 1, 1]]502    """503    ret = torch.zeros(size, size, device=device, dtype=torch.bool)504    for i in range(size):505        if num_left_chunks < 0:506            start = 0507        else:508            start = max((i // chunk_size - num_left_chunks) * chunk_size, 0)509        ending = min((i // chunk_size + 1) * chunk_size, size)510        ret[i, start:ending] = True511    return ret512 513def add_optional_chunk_mask(xs: torch.Tensor,514                            masks: torch.Tensor,515                            use_dynamic_chunk: bool,516                            use_dynamic_left_chunk: bool,517                            decoding_chunk_size: int,518                            static_chunk_size: int,519                            num_decoding_left_chunks: int,520                            enable_full_context: bool = True):521    """ Apply optional mask for encoder.522 523    Args:524        xs (torch.Tensor): padded input, (B, L, D), L for max length525        mask (torch.Tensor): mask for xs, (B, 1, L)526        use_dynamic_chunk (bool): whether to use dynamic chunk or not527        use_dynamic_left_chunk (bool): whether to use dynamic left chunk for528            training.529        decoding_chunk_size (int): decoding chunk size for dynamic chunk, it's530            0: default for training, use random dynamic chunk.531            <0: for decoding, use full chunk.532            >0: for decoding, use fixed chunk size as set.533        static_chunk_size (int): chunk size for static chunk training/decoding534            if it's greater than 0, if use_dynamic_chunk is true,535            this parameter will be ignored536        num_decoding_left_chunks: number of left chunks, this is for decoding,537            the chunk size is decoding_chunk_size.538            >=0: use num_decoding_left_chunks539            <0: use all left chunks540        enable_full_context (bool):541            True: chunk size is either [1, 25] or full context(max_len)542            False: chunk size ~ U[1, 25]543 544    Returns:545        torch.Tensor: chunk mask of the input xs.546    """547    # Whether to use chunk mask or not548    if use_dynamic_chunk:549        max_len = xs.size(1)550        if decoding_chunk_size < 0:551            chunk_size = max_len552            num_left_chunks = -1553        elif decoding_chunk_size > 0:554            chunk_size = decoding_chunk_size555            num_left_chunks = num_decoding_left_chunks556        else:557            # chunk size is either [1, 25] or full context(max_len).558            # Since we use 4 times subsampling and allow up to 1s(100 frames)559            # delay, the maximum frame is 100 / 4 = 25.560            chunk_size = torch.randint(1, max_len, (1, )).item()561            num_left_chunks = -1562            if chunk_size > max_len // 2 and enable_full_context:563                chunk_size = max_len564            else:565                chunk_size = chunk_size % 25 + 1566                if use_dynamic_left_chunk:567                    max_left_chunks = (max_len - 1) // chunk_size568                    num_left_chunks = torch.randint(0, max_left_chunks,569                                                    (1, )).item()570        chunk_masks = subsequent_chunk_mask(xs.size(1), chunk_size,571                                            num_left_chunks,572                                            xs.device)  # (L, L)573        chunk_masks = chunk_masks.unsqueeze(0)  # (1, L, L)574        chunk_masks = masks & chunk_masks  # (B, L, L)575    elif static_chunk_size > 0:576        num_left_chunks = num_decoding_left_chunks577        chunk_masks = subsequent_chunk_mask(xs.size(1), static_chunk_size,578                                            num_left_chunks,579                                            xs.device)  # (L, L)580        chunk_masks = chunk_masks.unsqueeze(0)  # (1, L, L)581        chunk_masks = masks & chunk_masks  # (B, L, L)582    else:583        chunk_masks = masks584    return chunk_masks585 586 587class ConformerEncoderLayer(nn.Module):588    """Encoder layer module.589    Args:590        size (int): Input dimension.591        self_attn (torch.nn.Module): Self-attention module instance.592            `MultiHeadedAttention` or `RelPositionMultiHeadedAttention`593            instance can be used as the argument.594        feed_forward (torch.nn.Module): Feed-forward module instance.595            `PositionwiseFeedForward` instance can be used as the argument.596        feed_forward_macaron (torch.nn.Module): Additional feed-forward module597             instance.598            `PositionwiseFeedForward` instance can be used as the argument.599        conv_module (torch.nn.Module): Convolution module instance.600            `ConvlutionModule` instance can be used as the argument.601        dropout_rate (float): Dropout rate.602        normalize_before (bool):603            True: use layer_norm before each sub-block.604            False: use layer_norm after each sub-block.605    """606 607    def __init__(608        self,609        size: int,610        self_attn: torch.nn.Module,611        feed_forward: Optional[nn.Module] = None,612        feed_forward_macaron: Optional[nn.Module] = None,613        conv_module: Optional[nn.Module] = None,614        dropout_rate: float = 0.1,615        normalize_before: bool = True,616    ):617        """Construct an EncoderLayer object."""618        super().__init__()619        self.self_attn = self_attn620        self.feed_forward = feed_forward621        self.feed_forward_macaron = feed_forward_macaron622        self.conv_module = conv_module623        self.norm_ff = nn.LayerNorm(size, eps=1e-5)  # for the FNN module624        self.norm_mha = nn.LayerNorm(size, eps=1e-5)  # for the MHA module625        if feed_forward_macaron is not None:626            self.norm_ff_macaron = nn.LayerNorm(size, eps=1e-5)627            self.ff_scale = 0.5628        else:629            self.ff_scale = 1.0630        if self.conv_module is not None:631            self.norm_conv = nn.LayerNorm(size, eps=1e-5)  # for the CNN module632            self.norm_final = nn.LayerNorm(633                size, eps=1e-5)  # for the final output of the block634        self.dropout = nn.Dropout(dropout_rate)635        self.size = size636        self.normalize_before = normalize_before637 638    def forward(639        self,640        x: torch.Tensor,641        mask: torch.Tensor,642        pos_emb: torch.Tensor,643        mask_pad: torch.Tensor = torch.ones((0, 0, 0), dtype=torch.bool),644        att_cache: torch.Tensor = torch.zeros((0, 0, 0, 0)),645        cnn_cache: torch.Tensor = torch.zeros((0, 0, 0, 0)),646    ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:647        """Compute encoded features.648 649        Args:650            x (torch.Tensor): (#batch, time, size)651            mask (torch.Tensor): Mask tensor for the input (#batch, time๏ผŒtime),652                (0, 0, 0) means fake mask.653            pos_emb (torch.Tensor): positional encoding, must not be None654                for ConformerEncoderLayer.655            mask_pad (torch.Tensor): batch padding mask used for conv module.656                (#batch, 1๏ผŒtime), (0, 0, 0) means fake mask.657            att_cache (torch.Tensor): Cache tensor of the KEY & VALUE658                (#batch=1, head, cache_t1, d_k * 2), head * d_k == size.659            cnn_cache (torch.Tensor): Convolution cache in conformer layer660                (#batch=1, size, cache_t2)661        Returns:662            torch.Tensor: Output tensor (#batch, time, size).663            torch.Tensor: Mask tensor (#batch, time, time).664            torch.Tensor: att_cache tensor,665                (#batch=1, head, cache_t1 + time, d_k * 2).666            torch.Tensor: cnn_cahce tensor (#batch, size, cache_t2).667        """668 669        # whether to use macaron style670        if self.feed_forward_macaron is not None:671            residual = x672            if self.normalize_before:673                x = self.norm_ff_macaron(x)674            x = residual + self.ff_scale * self.dropout(675                self.feed_forward_macaron(x))676            if not self.normalize_before:677                x = self.norm_ff_macaron(x)678 679        # multi-headed self-attention module680        residual = x681        if self.normalize_before:682            x = self.norm_mha(x)683        x_att, new_att_cache = self.self_attn(x, x, x, mask, pos_emb,684                                              att_cache)685        x = residual + self.dropout(x_att)686        if not self.normalize_before:687            x = self.norm_mha(x)688 689        # convolution module690        # Fake new cnn cache here, and then change it in conv_module691        new_cnn_cache = torch.zeros((0, 0, 0), dtype=x.dtype, device=x.device)692        if self.conv_module is not None:693            residual = x694            if self.normalize_before:695                x = self.norm_conv(x)696            x, new_cnn_cache = self.conv_module(x, mask_pad, cnn_cache)697            x = residual + self.dropout(x)698 699            if not self.normalize_before:700                x = self.norm_conv(x)701 702        # feed forward module703        residual = x704        if self.normalize_before:705            x = self.norm_ff(x)706 707        x = residual + self.ff_scale * self.dropout(self.feed_forward(x))708        if not self.normalize_before:709            x = self.norm_ff(x)710 711        if self.conv_module is not None:712            x = self.norm_final(x)713 714        return x, mask, new_att_cache, new_cnn_cache715    716 717 718class EspnetRelPositionalEncoding(torch.nn.Module):719    """Relative positional encoding module (new implementation).720 721    Details can be found in https://github.com/espnet/espnet/pull/2816.722 723    See : Appendix B in https://arxiv.org/abs/1901.02860724 725    Args:726        d_model (int): Embedding dimension.727        dropout_rate (float): Dropout rate.728        max_len (int): Maximum input length.729 730    """731 732    def __init__(self, d_model: int, dropout_rate: float, max_len: int = 5000):733        """Construct an PositionalEncoding object."""734        super(EspnetRelPositionalEncoding, self).__init__()735        self.d_model = d_model736        self.xscale = math.sqrt(self.d_model)737        self.dropout = torch.nn.Dropout(p=dropout_rate)738        self.pe = None739        self.extend_pe(torch.tensor(0.0).expand(1, max_len))740 741    def extend_pe(self, x: torch.Tensor):742        """Reset the positional encodings."""743        if self.pe is not None:744            # self.pe contains both positive and negative parts745            # the length of self.pe is 2 * input_len - 1746            if self.pe.size(1) >= x.size(1) * 2 - 1:747                if self.pe.dtype != x.dtype or self.pe.device != x.device:748                    self.pe = self.pe.to(dtype=x.dtype, device=x.device)749                return750        # Suppose `i` means to the position of query vecotr and `j` means the751        # position of key vector. We use position relative positions when keys752        # are to the left (i>j) and negative relative positions otherwise (i<j).753        pe_positive = torch.zeros(x.size(1), self.d_model)754        pe_negative = torch.zeros(x.size(1), self.d_model)755        position = torch.arange(0, x.size(1), dtype=torch.float32).unsqueeze(1)756        div_term = torch.exp(757            torch.arange(0, self.d_model, 2, dtype=torch.float32)758            * -(math.log(10000.0) / self.d_model)759        )760        pe_positive[:, 0::2] = torch.sin(position * div_term)761        pe_positive[:, 1::2] = torch.cos(position * div_term)762        pe_negative[:, 0::2] = torch.sin(-1 * position * div_term)763        pe_negative[:, 1::2] = torch.cos(-1 * position * div_term)764 765        # Reserve the order of positive indices and concat both positive and766        # negative indices. This is used to support the shifting trick767        # as in https://arxiv.org/abs/1901.02860768        pe_positive = torch.flip(pe_positive, [0]).unsqueeze(0)769        pe_negative = pe_negative[1:].unsqueeze(0)770        pe = torch.cat([pe_positive, pe_negative], dim=1)771        self.pe = pe.to(device=x.device, dtype=x.dtype)772 773    def forward(self, x: torch.Tensor, offset: Union[int, torch.Tensor] = 0) \774            -> Tuple[torch.Tensor, torch.Tensor]:775        """Add positional encoding.776 777        Args:778            x (torch.Tensor): Input tensor (batch, time, `*`).779 780        Returns:781            torch.Tensor: Encoded tensor (batch, time, `*`).782 783        """784        self.extend_pe(x)785        x = x * self.xscale786        pos_emb = self.position_encoding(size=x.size(1), offset=offset)787        return self.dropout(x), self.dropout(pos_emb)788    789    def position_encoding(self,790                          offset: Union[int, torch.Tensor],791                          size: int) -> torch.Tensor:792        """ For getting encoding in a streaming fashion793 794        Attention!!!!!795        we apply dropout only once at the whole utterance level in a none796        streaming way, but will call this function several times with797        increasing input size in a streaming scenario, so the dropout will798        be applied several times.799 800        Args:801            offset (int or torch.tensor): start offset802            size (int): required size of position encoding803 804        Returns:805            torch.Tensor: Corresponding encoding806        """807        pos_emb = self.pe[808            :,809            self.pe.size(1) // 2 - size + 1: self.pe.size(1) // 2 + size,810        ]811        return pos_emb812 813 814 815class LinearEmbed(torch.nn.Module):816    """Linear transform the input without subsampling817 818    Args:819        idim (int): Input dimension.820        odim (int): Output dimension.821        dropout_rate (float): Dropout rate.822 823    """824 825    def __init__(self, idim: int, odim: int, dropout_rate: float,826                 pos_enc_class: torch.nn.Module):827        """Construct an linear object."""828        super().__init__()829        self.out = torch.nn.Sequential(830            torch.nn.Linear(idim, odim),831            torch.nn.LayerNorm(odim, eps=1e-5),832            torch.nn.Dropout(dropout_rate),833        )834        self.pos_enc = pos_enc_class #rel_pos_espnet835    836    def position_encoding(self, offset: Union[int, torch.Tensor],837                          size: int) -> torch.Tensor:838        return self.pos_enc.position_encoding(offset, size)839 840    def forward(841        self,842        x: torch.Tensor,843        offset: Union[int, torch.Tensor] = 0844    ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:845        """Input x.846 847        Args:848            x (torch.Tensor): Input tensor (#batch, time, idim).849            x_mask (torch.Tensor): Input mask (#batch, 1, time).850 851        Returns:852            torch.Tensor: linear input tensor (#batch, time', odim),853                where time' = time .854            torch.Tensor: linear input mask (#batch, 1, time'),855                where time' = time .856 857        """858        x = self.out(x)859        x, pos_emb = self.pos_enc(x, offset)860        return x, pos_emb861 862 863ATTENTION_CLASSES = {864    "selfattn": MultiHeadedAttention,865    "rel_selfattn": RelPositionMultiHeadedAttention,866}867 868ACTIVATION_CLASSES = {869    "hardtanh": torch.nn.Hardtanh,870    "tanh": torch.nn.Tanh,871    "relu": torch.nn.ReLU,872    "selu": torch.nn.SELU,873    "swish": getattr(torch.nn, "SiLU", Swish),874    "gelu": torch.nn.GELU,875}876 877 878def make_pad_mask(lengths: torch.Tensor, max_len: int = 0) -> torch.Tensor:879    """Make mask tensor containing indices of padded part.880 881    See description of make_non_pad_mask.882 883    Args:884        lengths (torch.Tensor): Batch of lengths (B,).885    Returns:886        torch.Tensor: Mask tensor containing indices of padded part.887 888    Examples:889        >>> lengths = [5, 3, 2]890        >>> make_pad_mask(lengths)891        masks = [[0, 0, 0, 0 ,0],892                 [0, 0, 0, 1, 1],893                 [0, 0, 1, 1, 1]]894    """895    batch_size = lengths.size(0)896    max_len = max_len if max_len > 0 else lengths.max().item()897    seq_range = torch.arange(0,898                             max_len,899                             dtype=torch.int64,900                             device=lengths.device)901    seq_range_expand = seq_range.unsqueeze(0).expand(batch_size, max_len)902    seq_length_expand = lengths.unsqueeze(-1)903    mask = seq_range_expand >= seq_length_expand904    return mask905 906#https://github.com/FunAudioLLM/CosyVoice/blob/main/examples/magicdata-read/cosyvoice/conf/cosyvoice.yaml907class ConformerEncoder(torch.nn.Module):908    """Conformer encoder module."""909 910    def __init__(911        self,912        input_size: int,913        output_size: int = 1024,914        attention_heads: int = 16,915        linear_units: int = 4096,916        num_blocks: int = 6,917        dropout_rate: float = 0.1,918        positional_dropout_rate: float = 0.1,919        attention_dropout_rate: float = 0.0,920        input_layer: str = 'linear',921        pos_enc_layer_type: str = 'rel_pos_espnet',922        normalize_before: bool = True,923        static_chunk_size: int = 1, # 1: causal_mask; 0: full_mask924        use_dynamic_chunk: bool = False,925        use_dynamic_left_chunk: bool = False,926        positionwise_conv_kernel_size: int = 1,927        macaron_style: bool =False,928        selfattention_layer_type: str = "rel_selfattn",929        activation_type: str = "swish",930        use_cnn_module: bool = False,931        cnn_module_kernel: int = 15,932        causal: bool = False,933        cnn_module_norm: str = "batch_norm",934        key_bias: bool = True,935        gradient_checkpointing: bool = False,936    ):937        """Construct ConformerEncoder938 939        Args:940            input_size to use_dynamic_chunk, see in BaseEncoder941            positionwise_conv_kernel_size (int): Kernel size of positionwise942                conv1d layer.943            macaron_style (bool): Whether to use macaron style for944                positionwise layer.945            selfattention_layer_type (str): Encoder attention layer type,946                the parameter has no effect now, it's just for configure947                compatibility. #'rel_selfattn'948            activation_type (str): Encoder activation function type.949            use_cnn_module (bool): Whether to use convolution module.950            cnn_module_kernel (int): Kernel size of convolution module.951            causal (bool): whether to use causal convolution or not.952            key_bias: whether use bias in attention.linear_k, False for whisper models.953        """954        super().__init__()955        self.output_size = output_size956        self.embed = LinearEmbed(input_size, output_size, dropout_rate, 957                                        EspnetRelPositionalEncoding(output_size, positional_dropout_rate))958        self.normalize_before = normalize_before959        self.after_norm = torch.nn.LayerNorm(output_size, eps=1e-5)960        self.gradient_checkpointing = gradient_checkpointing961        self.use_dynamic_chunk = use_dynamic_chunk962        963        self.static_chunk_size = static_chunk_size964        self.use_dynamic_chunk = use_dynamic_chunk965        self.use_dynamic_left_chunk = use_dynamic_left_chunk966        activation = ACTIVATION_CLASSES[activation_type]()967 968        # self-attention module definition969        encoder_selfattn_layer_args = (970            attention_heads,971            output_size,972            attention_dropout_rate,973            key_bias,974        )975        # feed-forward module definition976        positionwise_layer_args = (977            output_size,978            linear_units,979            dropout_rate,980            activation,981        )982        # convolution module definition983        convolution_layer_args = (output_size, cnn_module_kernel, activation,984                                  cnn_module_norm, causal)985 986        self.encoders = torch.nn.ModuleList([987            ConformerEncoderLayer(988                output_size,989                RelPositionMultiHeadedAttention(990                    *encoder_selfattn_layer_args),991                PositionwiseFeedForward(*positionwise_layer_args),992                PositionwiseFeedForward(993                    *positionwise_layer_args) if macaron_style else None,994                ConvolutionModule(995                    *convolution_layer_args) if use_cnn_module else None,996                dropout_rate,997                normalize_before,998            ) for _ in range(num_blocks)999        ])1000    1001    def forward_layers(self, xs: torch.Tensor, chunk_masks: torch.Tensor,1002        pos_emb: torch.Tensor,1003        mask_pad: torch.Tensor) -> torch.Tensor:1004        for layer in self.encoders:1005            xs, chunk_masks, _, _ = layer(xs, chunk_masks, pos_emb, mask_pad)1006        return xs1007 1008    @torch.jit.unused1009    def forward_layers_checkpointed(self, xs: torch.Tensor,1010                                    chunk_masks: torch.Tensor,1011                                    pos_emb: torch.Tensor,1012                                    mask_pad: torch.Tensor) -> torch.Tensor:1013        for layer in self.encoders:1014            xs, chunk_masks, _, _ = ckpt.checkpoint(layer.__call__, xs,1015                                                    chunk_masks, pos_emb,1016                                                    mask_pad)1017        return xs1018 1019    def forward(1020        self,1021        xs: torch.Tensor,1022        pad_mask: torch.Tensor,1023        decoding_chunk_size: int = 0,1024        num_decoding_left_chunks: int = -1,1025    ) -> Tuple[torch.Tensor, torch.Tensor]:1026        """Embed positions in tensor.1027 1028        Args:1029            xs: padded input tensor (B, T, D)1030            xs_lens: input length (B)1031            decoding_chunk_size: decoding chunk size for dynamic chunk1032                0: default for training, use random dynamic chunk.1033                <0: for decoding, use full chunk.1034                >0: for decoding, use fixed chunk size as set.1035            num_decoding_left_chunks: number of left chunks, this is for decoding,1036            the chunk size is decoding_chunk_size.1037                >=0: use num_decoding_left_chunks1038                <0: use all left chunks1039        Returns:1040            encoder output tensor xs, and subsampled masks1041            xs: padded output tensor (B, T' ~= T/subsample_rate, D)1042            masks: torch.Tensor batch padding mask after subsample1043                (B, 1, T' ~= T/subsample_rate)1044        NOTE(xcsong):1045            We pass the `__call__` method of the modules instead of `forward` to the1046            checkpointing API because `__call__` attaches all the hooks of the module.1047            https://discuss.pytorch.org/t/any-different-between-model-input-and-model-forward-input/3690/21048        """1049        T = xs.size(1)1050        masks = pad_mask.to(torch.bool).unsqueeze(1)  # (B, 1, T) 1051        xs, pos_emb = self.embed(xs)1052        mask_pad = masks  # (B, 1, T/subsample_rate)1053        chunk_masks = add_optional_chunk_mask(xs, masks,1054                                              self.use_dynamic_chunk,1055                                              self.use_dynamic_left_chunk,1056                                              decoding_chunk_size,1057                                              self.static_chunk_size,1058                                              num_decoding_left_chunks) 1059        if self.gradient_checkpointing and self.training:1060            xs = self.forward_layers_checkpointed(xs, chunk_masks, pos_emb,1061                                                  mask_pad)1062        else:1063            xs = self.forward_layers(xs, chunk_masks, pos_emb, mask_pad)1064        if self.normalize_before:1065            xs = self.after_norm(xs)1066        # Here we assume the mask is not changed in encoder layers, so just1067        # return the masks before encoder layers, and the masks will be used1068        # for cross attention with decoder later1069        return xs, masks1070 1071