MCPcopy Create free account
hub / github.com/OpenImagingLab/FlashVSR / IFBlock

Class IFBlock

diffsynth/extensions/RIFE/__init__.py:34–57  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

32
33
34class IFBlock(nn.Module):
35 def __init__(self, in_planes, c=64):
36 super(IFBlock, self).__init__()
37 self.conv0 = nn.Sequential(conv(in_planes, c//2, 3, 2, 1), conv(c//2, c, 3, 2, 1),)
38 self.convblock0 = nn.Sequential(conv(c, c), conv(c, c))
39 self.convblock1 = nn.Sequential(conv(c, c), conv(c, c))
40 self.convblock2 = nn.Sequential(conv(c, c), conv(c, c))
41 self.convblock3 = nn.Sequential(conv(c, c), conv(c, c))
42 self.conv1 = nn.Sequential(nn.ConvTranspose2d(c, c//2, 4, 2, 1), nn.PReLU(c//2), nn.ConvTranspose2d(c//2, 4, 4, 2, 1))
43 self.conv2 = nn.Sequential(nn.ConvTranspose2d(c, c//2, 4, 2, 1), nn.PReLU(c//2), nn.ConvTranspose2d(c//2, 1, 4, 2, 1))
44
45 def forward(self, x, flow, scale=1):
46 x = F.interpolate(x, scale_factor= 1. / scale, mode="bilinear", align_corners=False, recompute_scale_factor=False)
47 flow = F.interpolate(flow, scale_factor= 1. / scale, mode="bilinear", align_corners=False, recompute_scale_factor=False) * 1. / scale
48 feat = self.conv0(torch.cat((x, flow), 1))
49 feat = self.convblock0(feat) + feat
50 feat = self.convblock1(feat) + feat
51 feat = self.convblock2(feat) + feat
52 feat = self.convblock3(feat) + feat
53 flow = self.conv1(feat)
54 mask = self.conv2(feat)
55 flow = F.interpolate(flow, scale_factor=scale, mode="bilinear", align_corners=False, recompute_scale_factor=False) * scale
56 mask = F.interpolate(mask, scale_factor=scale, mode="bilinear", align_corners=False, recompute_scale_factor=False)
57 return flow, mask
58
59
60class IFNet(nn.Module):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected