Skip to content

ZeroInit #8

Description

@ibro45

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

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)

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions