MCPcopy Create free account
hub / github.com/UVA-Computer-Vision-Lab/FrameINO / __init__

Method __init__

architecture/autoencoder_kl_wan.py:234–263  ·  view source on GitHub ↗
(self, dim: int, mode: str, upsample_out_dim: int = None)

Source from the content-addressed store, hash-verified

232 """
233
234 def __init__(self, dim: int, mode: str, upsample_out_dim: int = None) -> None:
235 super().__init__()
236 self.dim = dim
237 self.mode = mode
238
239 # default to dim //2
240 if upsample_out_dim is None:
241 upsample_out_dim = dim // 2
242
243 # layers
244 if mode == "upsample2d":
245 self.resample = nn.Sequential(
246 WanUpsample(scale_factor=(2.0, 2.0), mode="nearest-exact"),
247 nn.Conv2d(dim, upsample_out_dim, 3, padding=1),
248 )
249 elif mode == "upsample3d":
250 self.resample = nn.Sequential(
251 WanUpsample(scale_factor=(2.0, 2.0), mode="nearest-exact"),
252 nn.Conv2d(dim, upsample_out_dim, 3, padding=1),
253 )
254 self.time_conv = WanCausalConv3d(dim, dim * 2, (3, 1, 1), padding=(1, 0, 0))
255
256 elif mode == "downsample2d":
257 self.resample = nn.Sequential(nn.ZeroPad2d((0, 1, 0, 1)), nn.Conv2d(dim, dim, 3, stride=(2, 2)))
258 elif mode == "downsample3d":
259 self.resample = nn.Sequential(nn.ZeroPad2d((0, 1, 0, 1)), nn.Conv2d(dim, dim, 3, stride=(2, 2)))
260 self.time_conv = WanCausalConv3d(dim, dim, (3, 1, 1), stride=(2, 1, 1), padding=(0, 0, 0))
261
262 else:
263 self.resample = nn.Identity()
264
265 def forward(self, x, feat_cache=None, feat_idx=[0]):
266 b, c, t, h, w = x.size()

Callers

nothing calls this directly

Calls 3

WanUpsampleClass · 0.85
WanCausalConv3dClass · 0.85
__init__Method · 0.45

Tested by

no test coverage detected