Construct an DecoderLayer object.
(
self,
size: int,
self_attn: nn.Module,
src_attn: Optional[nn.Module],
feed_forward: nn.Module,
dropout_rate: float,
normalize_before: bool = True,
layer_norm_type: str = 'layer_norm',
norm_eps: float = 1e-5,
)
| 42 | """ |
| 43 | |
| 44 | def __init__( |
| 45 | self, |
| 46 | size: int, |
| 47 | self_attn: nn.Module, |
| 48 | src_attn: Optional[nn.Module], |
| 49 | feed_forward: nn.Module, |
| 50 | dropout_rate: float, |
| 51 | normalize_before: bool = True, |
| 52 | layer_norm_type: str = 'layer_norm', |
| 53 | norm_eps: float = 1e-5, |
| 54 | ): |
| 55 | """Construct an DecoderLayer object.""" |
| 56 | super().__init__() |
| 57 | self.size = size |
| 58 | self.self_attn = self_attn |
| 59 | self.src_attn = src_attn |
| 60 | self.feed_forward = feed_forward |
| 61 | assert layer_norm_type in ['layer_norm', 'rms_norm'] |
| 62 | self.norm1 = WENET_NORM_CLASSES[layer_norm_type](size, eps=norm_eps) |
| 63 | self.norm2 = WENET_NORM_CLASSES[layer_norm_type](size, eps=norm_eps) |
| 64 | self.norm3 = WENET_NORM_CLASSES[layer_norm_type](size, eps=norm_eps) |
| 65 | self.dropout = nn.Dropout(dropout_rate) |
| 66 | self.normalize_before = normalize_before |
| 67 | |
| 68 | def forward( |
| 69 | self, |
nothing calls this directly
no outgoing calls
no test coverage detected