MCPcopy Create free account
hub / github.com/FunAudioLLM/FunMusic / TransformerEncoderLayer

Class TransformerEncoderLayer

inspiremusic/transformer/encoder_layer.py:24–106  ·  view source on GitHub ↗

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):

Source from the content-addressed store, hash-verified

22
23
24class TransformerEncoderLayer(nn.Module):
25 """Encoder layer module.
26
27 Args:
28 size (int): Input dimension.
29 self_attn (torch.nn.Module): Self-attention module instance.
30 `MultiHeadedAttention` or `RelPositionMultiHeadedAttention`
31 instance can be used as the argument.
32 feed_forward (torch.nn.Module): Feed-forward module instance.
33 `PositionwiseFeedForward`, instance can be used as the argument.
34 dropout_rate (float): Dropout rate.
35 normalize_before (bool):
36 True: use layer_norm before each sub-block.
37 False: to use layer_norm after each sub-block.
38 """
39
40 def __init__(
41 self,
42 size: int,
43 self_attn: torch.nn.Module,
44 feed_forward: torch.nn.Module,
45 dropout_rate: float,
46 normalize_before: bool = True,
47 ):
48 """Construct an EncoderLayer object."""
49 super().__init__()
50 self.self_attn = self_attn
51 self.feed_forward = feed_forward
52 self.norm1 = nn.LayerNorm(size, eps=1e-5)
53 self.norm2 = nn.LayerNorm(size, eps=1e-5)
54 self.dropout = nn.Dropout(dropout_rate)
55 self.size = size
56 self.normalize_before = normalize_before
57
58 def forward(
59 self,
60 x: torch.Tensor,
61 mask: torch.Tensor,
62 pos_emb: torch.Tensor,
63 mask_pad: torch.Tensor = torch.ones((0, 0, 0), dtype=torch.bool),
64 att_cache: torch.Tensor = torch.zeros((0, 0, 0, 0)),
65 cnn_cache: torch.Tensor = torch.zeros((0, 0, 0, 0)),
66 ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
67 """Compute encoded features.
68
69 Args:
70 x (torch.Tensor): (#batch, time, size)
71 mask (torch.Tensor): Mask tensor for the input (#batch, time,time),
72 (0, 0, 0) means fake mask.
73 pos_emb (torch.Tensor): just for interface compatibility
74 to ConformerEncoderLayer
75 mask_pad (torch.Tensor): does not used in transformer layer,
76 just for unified api with conformer.
77 att_cache (torch.Tensor): Cache tensor of the KEY & VALUE
78 (#batch=1, head, cache_t1, d_k * 2), head * d_k == size.
79 cnn_cache (torch.Tensor): Convolution cache in conformer layer
80 (#batch=1, size, cache_t2), not used here, it's for interface
81 compatibility to ConformerEncoderLayer.

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected