| 292 | |
| 293 | |
| 294 | class EncoderLayerSANM(nn.Module): |
| 295 | def __init__( |
| 296 | self, |
| 297 | in_size, |
| 298 | size, |
| 299 | self_attn, |
| 300 | feed_forward, |
| 301 | dropout_rate, |
| 302 | normalize_before=True, |
| 303 | concat_after=False, |
| 304 | stochastic_depth_rate=0.0, |
| 305 | ): |
| 306 | """Construct an EncoderLayer object.""" |
| 307 | super(EncoderLayerSANM, self).__init__() |
| 308 | self.self_attn = self_attn |
| 309 | self.feed_forward = feed_forward |
| 310 | self.norm1 = LayerNorm(in_size) |
| 311 | self.norm2 = LayerNorm(size) |
| 312 | self.dropout = nn.Dropout(dropout_rate) |
| 313 | self.in_size = in_size |
| 314 | self.size = size |
| 315 | self.normalize_before = normalize_before |
| 316 | self.concat_after = concat_after |
| 317 | if self.concat_after: |
| 318 | self.concat_linear = nn.Linear(size + size, size) |
| 319 | self.stochastic_depth_rate = stochastic_depth_rate |
| 320 | self.dropout_rate = dropout_rate |
| 321 | |
| 322 | def forward(self, x, mask, cache=None, mask_shfit_chunk=None, mask_att_chunk_encoder=None): |
| 323 | """Compute encoded features. |
| 324 | |
| 325 | Args: |
| 326 | x_input (torch.Tensor): Input tensor (#batch, time, size). |
| 327 | mask (torch.Tensor): Mask tensor for the input (#batch, time). |
| 328 | cache (torch.Tensor): Cache tensor of the input (#batch, time - 1, size). |
| 329 | |
| 330 | Returns: |
| 331 | torch.Tensor: Output tensor (#batch, time, size). |
| 332 | torch.Tensor: Mask tensor (#batch, time). |
| 333 | |
| 334 | """ |
| 335 | skip_layer = False |
| 336 | # with stochastic depth, residual connection `x + f(x)` becomes |
| 337 | # `x <- x + 1 / (1 - p) * f(x)` at training time. |
| 338 | stoch_layer_coeff = 1.0 |
| 339 | if self.training and self.stochastic_depth_rate > 0: |
| 340 | skip_layer = torch.rand(1).item() < self.stochastic_depth_rate |
| 341 | stoch_layer_coeff = 1.0 / (1 - self.stochastic_depth_rate) |
| 342 | |
| 343 | if skip_layer: |
| 344 | if cache is not None: |
| 345 | x = torch.cat([cache, x], dim=1) |
| 346 | return x, mask |
| 347 | |
| 348 | residual = x |
| 349 | if self.normalize_before: |
| 350 | x = self.norm1(x) |
| 351 | |