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

Method forward

diffsynth/models/cog_vae.py:111–124  ·  view source on GitHub ↗
(self, f: torch.Tensor, zq: torch.Tensor)

Source from the content-addressed store, hash-verified

109
110
111 def forward(self, f: torch.Tensor, zq: torch.Tensor) -> torch.Tensor:
112 if f.shape[2] > 1 and f.shape[2] % 2 == 1:
113 f_first, f_rest = f[:, :, :1], f[:, :, 1:]
114 f_first_size, f_rest_size = f_first.shape[-3:], f_rest.shape[-3:]
115 z_first, z_rest = zq[:, :, :1], zq[:, :, 1:]
116 z_first = torch.nn.functional.interpolate(z_first, size=f_first_size)
117 z_rest = torch.nn.functional.interpolate(z_rest, size=f_rest_size)
118 zq = torch.cat([z_first, z_rest], dim=2)
119 else:
120 zq = torch.nn.functional.interpolate(zq, size=f.shape[-3:])
121
122 norm_f = self.norm_layer(f)
123 new_f = norm_f * self.conv_y(zq) + self.conv_b(zq)
124 return new_f
125
126
127

Callers

nothing calls this directly

Calls 1

interpolateMethod · 0.80

Tested by

no test coverage detected