|
class ZeroInit(nn.Conv3d): |
|
def reset_parameters(self): |
|
self.weight.data.zero_() |
|
self.bias.data.zero_() |
|
|
|
class ThickBlocknaive3d(nn.Module): |
|
def __init__(self, dim, use_bias): |
|
super(ThickBlocknaive3d, self).__init__() |
|
self.F = self.build_conv_block(dim, True) |
|
|
|
def build_conv_block(self, dim, use_bias): |
|
conv_block = [] |
|
conv_block += [nn.InstanceNorm3d(dim)] |
|
conv_block += [nn.ReplicationPad3d(1)] |
|
conv_block += [nn.Conv3d(dim, dim, kernel_size=3, padding=0, bias=use_bias)] |
|
conv_block += [nn.InstanceNorm3d(dim)] |
|
conv_block += [nn.ReLU(True)] |
|
conv_block += [nn.ReplicationPad3d(1)] |
|
conv_block += [ZeroInit(dim, dim, kernel_size=3, padding=0, bias=use_bias)] |
|
|
|
|
|
return nn.Sequential(*conv_block) |
|
|
|
def forward(self, input): |
|
return self.F(input) + input |
|
|
|
|
|
class ThickBlock3d(nn.Module): |
|
def __init__(self, dim, use_bias, use_naive=False): |
|
super(ThickBlock3d, self).__init__() |
|
F = self.build_conv_block(dim // 2, True) |
|
G = self.build_conv_block(dim // 2, True) |
|
if use_naive: |
|
self.rev_block = ReversibleBlock(F, G, 'additive', |
|
keep_input=True, implementation_fwd=2, implementation_bwd=2) |
|
else: |
|
self.rev_block = ReversibleBlock(F, G, 'additive') |
|
|
|
|
|
def build_conv_block(self, dim, use_bias): |
|
conv_block = [] |
|
conv_block += [nn.InstanceNorm3d(dim)] |
|
conv_block += [nn.ReplicationPad3d(1)] |
|
conv_block += [nn.Conv3d(dim, dim, kernel_size=3, padding=0, bias=use_bias)] |
|
conv_block += [nn.InstanceNorm3d(dim)] |
|
conv_block += [nn.ReLU(True)] |
|
conv_block += [nn.ReplicationPad3d(1)] |
|
conv_block += [ZeroInit(dim, dim, kernel_size=3, padding=0, bias=use_bias)] |
|
|
|
|
|
return nn.Sequential(*conv_block) |
|
|
|
def forward(self, x): |
|
return self.rev_block(x) |
|
|
|
def inverse(self, x): |
|
return self.rev_block.inverse(x) |
|
|
|
|
|
|
|
class RevBlock3d(nn.Module): |
|
def __init__(self, dim, use_bias, norm_layer, use_naive): |
|
super(RevBlock3d, self).__init__() |
|
self.F = self.build_conv_block(dim // 2, True, norm_layer) |
|
self.G = self.build_conv_block(dim // 2, True, norm_layer) |
|
if use_naive: |
|
self.rev_block = ReversibleBlock(F, G, 'additive', |
|
keep_input=True, implementation_fwd=2, implementation_bwd=2) |
|
else: |
|
self.rev_block = ReversibleBlock(F, G, 'additive') |
|
|
|
def build_conv_block(self, dim, use_bias, norm_layer): |
|
conv_block = [] |
|
conv_block += [nn.ReplicationPad3d(1)] |
|
conv_block += [nn.Conv3d(dim, dim, kernel_size=3, padding=0, bias=use_bias)] |
|
conv_block += [norm_layer(dim)] |
|
conv_block += [nn.ReLU(True)] |
|
conv_block += [nn.ReplicationPad3d(1)] |
|
conv_block += [ZeroInit(dim, dim, kernel_size=3, padding=0, bias=use_bias)] |
|
|
|
|
|
return nn.Sequential(*conv_block) |
|
|
|
def forward(self, x): |
|
return self.rev_block(x) |
|
|
|
def inverse(self, x): |
|
return self.rev_block.inverse(x) |
|
|
Hi! Could you please explain what is the purpose of ZeroInit in reversible blocks? Thanks!
RevGAN/models/networks3d.py
Lines 233 to 321 in 2af25e6