MCPcopy Create free account
hub / github.com/Francis-Rings/FlashPortrait / __init__

Method __init__

wan/models/wan_vae3_8.py:78–115  ·  view source on GitHub ↗
(self, dim, mode)

Source from the content-addressed store, hash-verified

76class 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()

Callers

nothing calls this directly

Calls 3

UpsampleClass · 0.70
CausalConv3dClass · 0.70
__init__Method · 0.45

Tested by

no test coverage detected