MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / __init__

Method __init__

sat/sgm/modules/autoencoding/vqvae/movq_modules.py:111–139  ·  view source on GitHub ↗
(
        self,
        *,
        in_channels,
        out_channels=None,
        conv_shortcut=False,
        dropout,
        temb_channels=512,
        zq_ch=None,
        add_conv=False,
    )

Source from the content-addressed store, hash-verified

109
110class ResnetBlock(nn.Module):
111 def __init__(
112 self,
113 *,
114 in_channels,
115 out_channels=None,
116 conv_shortcut=False,
117 dropout,
118 temb_channels=512,
119 zq_ch=None,
120 add_conv=False,
121 ):
122 super().__init__()
123 self.in_channels = in_channels
124 out_channels = in_channels if out_channels is None else out_channels
125 self.out_channels = out_channels
126 self.use_conv_shortcut = conv_shortcut
127
128 self.norm1 = Normalize(in_channels, zq_ch, add_conv=add_conv)
129 self.conv1 = torch.nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
130 if temb_channels > 0:
131 self.temb_proj = torch.nn.Linear(temb_channels, out_channels)
132 self.norm2 = Normalize(out_channels, zq_ch, add_conv=add_conv)
133 self.dropout = torch.nn.Dropout(dropout)
134 self.conv2 = torch.nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1)
135 if self.in_channels != self.out_channels:
136 if self.use_conv_shortcut:
137 self.conv_shortcut = torch.nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
138 else:
139 self.nin_shortcut = torch.nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, padding=0)
140
141 def forward(self, x, temb, zq):
142 h = x

Callers

nothing calls this directly

Calls 2

NormalizeFunction · 0.70
__init__Method · 0.45

Tested by

no test coverage detected