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

Class AvgDown3D

architecture/autoencoder_kl_wan.py:37–87  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

35
36
37class AvgDown3D(nn.Module):
38 def __init__(
39 self,
40 in_channels,
41 out_channels,
42 factor_t,
43 factor_s=1,
44 ):
45 super().__init__()
46 self.in_channels = in_channels
47 self.out_channels = out_channels
48 self.factor_t = factor_t
49 self.factor_s = factor_s
50 self.factor = self.factor_t * self.factor_s * self.factor_s
51
52 assert in_channels * self.factor % out_channels == 0
53 self.group_size = in_channels * self.factor // out_channels
54
55 def forward(self, x: torch.Tensor) -> torch.Tensor:
56 pad_t = (self.factor_t - x.shape[2] % self.factor_t) % self.factor_t
57 pad = (0, 0, 0, 0, pad_t, 0)
58 x = F.pad(x, pad)
59 B, C, T, H, W = x.shape
60 x = x.view(
61 B,
62 C,
63 T // self.factor_t,
64 self.factor_t,
65 H // self.factor_s,
66 self.factor_s,
67 W // self.factor_s,
68 self.factor_s,
69 )
70 x = x.permute(0, 1, 3, 5, 7, 2, 4, 6).contiguous()
71 x = x.view(
72 B,
73 C * self.factor,
74 T // self.factor_t,
75 H // self.factor_s,
76 W // self.factor_s,
77 )
78 x = x.view(
79 B,
80 self.out_channels,
81 self.group_size,
82 T // self.factor_t,
83 H // self.factor_s,
84 W // self.factor_s,
85 )
86 x = x.mean(dim=2)
87 return x
88
89
90class DupUp3D(nn.Module):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected