(self, dim, mode)
| 76 | class Resample(nn.Module): |
| 77 | |
| 78 | def __init__(self, dim, mode): |
| 79 | assert mode in ( |
| 80 | "none", |
| 81 | "upsample2d", |
| 82 | "upsample3d", |
| 83 | "downsample2d", |
| 84 | "downsample3d", |
| 85 | ) |
| 86 | super().__init__() |
| 87 | self.dim = dim |
| 88 | self.mode = mode |
| 89 | |
| 90 | # layers |
| 91 | if mode == "upsample2d": |
| 92 | self.resample = nn.Sequential( |
| 93 | Upsample(scale_factor=(2.0, 2.0), mode="nearest-exact"), |
| 94 | nn.Conv2d(dim, dim, 3, padding=1), |
| 95 | ) |
| 96 | elif mode == "upsample3d": |
| 97 | self.resample = nn.Sequential( |
| 98 | Upsample(scale_factor=(2.0, 2.0), mode="nearest-exact"), |
| 99 | nn.Conv2d(dim, dim, 3, padding=1), |
| 100 | # nn.Conv2d(dim, dim//2, 3, padding=1) |
| 101 | ) |
| 102 | self.time_conv = CausalConv3d( |
| 103 | dim, dim * 2, (3, 1, 1), padding=(1, 0, 0)) |
| 104 | elif mode == "downsample2d": |
| 105 | self.resample = nn.Sequential( |
| 106 | nn.ZeroPad2d((0, 1, 0, 1)), |
| 107 | nn.Conv2d(dim, dim, 3, stride=(2, 2))) |
| 108 | elif mode == "downsample3d": |
| 109 | self.resample = nn.Sequential( |
| 110 | nn.ZeroPad2d((0, 1, 0, 1)), |
| 111 | nn.Conv2d(dim, dim, 3, stride=(2, 2))) |
| 112 | self.time_conv = CausalConv3d( |
| 113 | dim, dim, (3, 1, 1), stride=(2, 1, 1), padding=(0, 0, 0)) |
| 114 | else: |
| 115 | self.resample = nn.Identity() |
| 116 | |
| 117 | def forward(self, x, feat_cache=None, feat_idx=[0]): |
| 118 | b, c, t, h, w = x.size() |
nothing calls this directly
no test coverage detected