MCPcopy Create free account
hub / github.com/kwuking/TimeMixer / Inception_Trans_Block_V1

Class Inception_Trans_Block_V1

layers/Conv_Blocks.py:34–62  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

32
33
34class Inception_Trans_Block_V1(nn.Module):
35 def __init__(self, in_channels, out_channels, stride=1, num_kernels=6, init_weight=True):
36 super(Inception_Trans_Block_V1, self).__init__()
37 self.in_channels = in_channels
38 self.out_channels = out_channels
39 self.num_kernels = num_kernels
40 self.stride = stride
41
42 kernels = []
43 for i in range(self.num_kernels):
44 kernels.append(
45 nn.ConvTranspose2d(in_channels, out_channels, kernel_size=2 * i + 1, padding=i, stride=stride))
46 self.kernels = nn.ModuleList(kernels)
47 if init_weight:
48 self._initialize_weights()
49
50 def _initialize_weights(self):
51 for m in self.modules():
52 if isinstance(m, nn.ConvTranspose2d):
53 nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')
54 if m.bias is not None:
55 nn.init.constant_(m.bias, 0)
56
57 def forward(self, x, output_size):
58 res_list = []
59 for i in range(self.num_kernels):
60 res_list.append(self.kernels[i](x, output_size=output_size))
61 res = torch.stack(res_list, dim=-1).mean(-1)
62 return res

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected