MCPcopy Create free account
hub / github.com/MoonInTheRiver/DiffSinger / TransformerFFNLayer

Class TransformerFFNLayer

modules/commons/common_layers.py:486–522  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

484
485
486class TransformerFFNLayer(nn.Module):
487 def __init__(self, hidden_size, filter_size, padding="SAME", kernel_size=1, dropout=0., act='gelu'):
488 super().__init__()
489 self.kernel_size = kernel_size
490 self.dropout = dropout
491 self.act = act
492 if padding == 'SAME':
493 self.ffn_1 = nn.Conv1d(hidden_size, filter_size, kernel_size, padding=kernel_size // 2)
494 elif padding == 'LEFT':
495 self.ffn_1 = nn.Sequential(
496 nn.ConstantPad1d((kernel_size - 1, 0), 0.0),
497 nn.Conv1d(hidden_size, filter_size, kernel_size)
498 )
499 self.ffn_2 = Linear(filter_size, hidden_size)
500 if self.act == 'swish':
501 self.swish_fn = CustomSwish()
502
503 def forward(self, x, incremental_state=None):
504 # x: T x B x C
505 if incremental_state is not None:
506 assert incremental_state is None, 'Nar-generation does not allow this.'
507 exit(1)
508
509 x = self.ffn_1(x.permute(1, 2, 0)).permute(2, 0, 1)
510 x = x * self.kernel_size ** -0.5
511
512 if incremental_state is not None:
513 x = x[-1:]
514 if self.act == 'gelu':
515 x = F.gelu(x)
516 if self.act == 'relu':
517 x = F.relu(x)
518 if self.act == 'swish':
519 x = self.swish_fn(x)
520 x = F.dropout(x, self.dropout, training=self.training)
521 x = self.ffn_2(x)
522 return x
523
524
525class BatchNorm1dTBC(nn.Module):

Callers 2

__init__Method · 0.85
__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected