KUI71/ACE-Step
0
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 