MCPcopy Create free account
hub / github.com/QWTforGithub/T2LDM / Downsample

Class Downsample

models/T2LDM.py:1191–1221  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

1189 return x
1190
1191class Downsample(nn.Module):
1192 def __init__(
1193 self,
1194 in_channels,
1195 stride,
1196 out_channels=-1,
1197 with_conv=False
1198 ):
1199 super().__init__()
1200 self.with_conv = with_conv
1201 self.stride = stride
1202 if(out_channels == -1):
1203 out_channels = in_channels
1204 if(self.with_conv == "CircularConv2D"):
1205 k, p = DOWNSAMPLE_STRIDE2KERNEL_DICT[stride], DOWNSAMPLE_STRIDE2PAD_DICT[stride]
1206 self.conv = CircularConv2D(in_channels, out_channels, kernel_size=k, stride=stride, padding=p)
1207 elif(self.with_conv == "Conv2D"):
1208 self.conv = nn.Conv2d(
1209 in_channels,
1210 out_channels,
1211 kernel_size=stride,
1212 stride=stride,
1213 padding=0
1214 )
1215
1216 def forward(self, x):
1217 if self.with_conv:
1218 x = self.conv(x)
1219 else:
1220 x = torch.nn.functional.avg_pool2d(x, kernel_size=self.stride, stride=self.stride)
1221 return x
1222
1223class Upsample(nn.Module):
1224 def __init__(

Callers 1

get_moduleFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected