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

Class ResnetBlock3D

sat/sgm/modules/autoencoding/vqvae/movq_dec_3d_dev.py:108–176  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

106
107
108class ResnetBlock3D(nn.Module):
109 def __init__(
110 self,
111 *,
112 in_channels,
113 out_channels=None,
114 conv_shortcut=False,
115 dropout,
116 temb_channels=512,
117 zq_ch=None,
118 add_conv=False,
119 pad_mode="constant",
120 ):
121 super().__init__()
122 self.in_channels = in_channels
123 out_channels = in_channels if out_channels is None else out_channels
124 self.out_channels = out_channels
125 self.use_conv_shortcut = conv_shortcut
126
127 self.norm1 = Normalize3D(in_channels, zq_ch, add_conv=add_conv)
128 # self.conv1 = torch.nn.Conv3d(in_channels,
129 # out_channels,
130 # kernel_size=3,
131 # stride=1,
132 # padding=1)
133 self.conv1 = CausalConv3d(in_channels, out_channels, kernel_size=3, pad_mode=pad_mode)
134 if temb_channels > 0:
135 self.temb_proj = torch.nn.Linear(temb_channels, out_channels)
136 self.norm2 = Normalize3D(out_channels, zq_ch, add_conv=add_conv)
137 self.dropout = torch.nn.Dropout(dropout)
138 # self.conv2 = torch.nn.Conv3d(out_channels,
139 # out_channels,
140 # kernel_size=3,
141 # stride=1,
142 # padding=1)
143 self.conv2 = CausalConv3d(out_channels, out_channels, kernel_size=3, pad_mode=pad_mode)
144 if self.in_channels != self.out_channels:
145 if self.use_conv_shortcut:
146 # self.conv_shortcut = torch.nn.Conv3d(in_channels,
147 # out_channels,
148 # kernel_size=3,
149 # stride=1,
150 # padding=1)
151 self.conv_shortcut = CausalConv3d(in_channels, out_channels, kernel_size=3, pad_mode=pad_mode)
152 else:
153 self.nin_shortcut = torch.nn.Conv3d(in_channels, out_channels, kernel_size=1, stride=1, padding=0)
154 # self.nin_shortcut = CausalConv3d(in_channels, out_channels, kernel_size=1, pad_mode=pad_mode)
155
156 def forward(self, x, temb, zq):
157 h = x
158 h = self.norm1(h, zq)
159 h = nonlinearity(h)
160 h = self.conv1(h)
161
162 if temb is not None:
163 h = h + self.temb_proj(nonlinearity(temb))[:, :, None, None, None]
164
165 h = self.norm2(h, zq)

Callers 2

__init__Method · 0.70
__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected