(self, dim: int, mode: str, upsample_out_dim: int = None)
| 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() |
nothing calls this directly
no test coverage detected