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

Class ResidualStack

modules/parallel_wavegan/layers/residual_stack.py:13–75  ·  view source on GitHub ↗

Residual stack module introduced in MelGAN.

Source from the content-addressed store, hash-verified

11
12
13class ResidualStack(torch.nn.Module):
14 """Residual stack module introduced in MelGAN."""
15
16 def __init__(self,
17 kernel_size=3,
18 channels=32,
19 dilation=1,
20 bias=True,
21 nonlinear_activation="LeakyReLU",
22 nonlinear_activation_params={"negative_slope": 0.2},
23 pad="ReflectionPad1d",
24 pad_params={},
25 use_causal_conv=False,
26 ):
27 """Initialize ResidualStack module.
28
29 Args:
30 kernel_size (int): Kernel size of dilation convolution layer.
31 channels (int): Number of channels of convolution layers.
32 dilation (int): Dilation factor.
33 bias (bool): Whether to add bias parameter in convolution layers.
34 nonlinear_activation (str): Activation function module name.
35 nonlinear_activation_params (dict): Hyperparameters for activation function.
36 pad (str): Padding function module name before dilated convolution layer.
37 pad_params (dict): Hyperparameters for padding function.
38 use_causal_conv (bool): Whether to use causal convolution.
39
40 """
41 super(ResidualStack, self).__init__()
42
43 # defile residual stack part
44 if not use_causal_conv:
45 assert (kernel_size - 1) % 2 == 0, "Not support even number kernel size."
46 self.stack = torch.nn.Sequential(
47 getattr(torch.nn, nonlinear_activation)(**nonlinear_activation_params),
48 getattr(torch.nn, pad)((kernel_size - 1) // 2 * dilation, **pad_params),
49 torch.nn.Conv1d(channels, channels, kernel_size, dilation=dilation, bias=bias),
50 getattr(torch.nn, nonlinear_activation)(**nonlinear_activation_params),
51 torch.nn.Conv1d(channels, channels, 1, bias=bias),
52 )
53 else:
54 self.stack = torch.nn.Sequential(
55 getattr(torch.nn, nonlinear_activation)(**nonlinear_activation_params),
56 CausalConv1d(channels, channels, kernel_size, dilation=dilation,
57 bias=bias, pad=pad, pad_params=pad_params),
58 getattr(torch.nn, nonlinear_activation)(**nonlinear_activation_params),
59 torch.nn.Conv1d(channels, channels, 1, bias=bias),
60 )
61
62 # defile extra layer for skip connection
63 self.skip_layer = torch.nn.Conv1d(channels, channels, 1, bias=bias)
64
65 def forward(self, c):
66 """Calculate forward propagation.
67
68 Args:
69 c (Tensor): Input tensor (B, channels, T).
70

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected