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

Function pos_interpolate

diffsynth/models/wan_video_image_encoder.py:203–219  ·  view source on GitHub ↗
(pos, seq_len)

Source from the content-addressed store, hash-verified

201
202
203def pos_interpolate(pos, seq_len):
204 if pos.size(1) == seq_len:
205 return pos
206 else:
207 src_grid = int(math.sqrt(pos.size(1)))
208 tar_grid = int(math.sqrt(seq_len))
209 n = pos.size(1) - src_grid * src_grid
210 return torch.cat([
211 pos[:, :n],
212 F.interpolate(
213 pos[:, n:].float().reshape(1, src_grid, src_grid, -1).permute(
214 0, 3, 1, 2),
215 size=(tar_grid, tar_grid),
216 mode='bicubic',
217 align_corners=False).flatten(2).transpose(1, 2)
218 ],
219 dim=1)
220
221
222class QuickGELU(nn.Module):

Callers 1

forwardMethod · 0.85

Calls 1

interpolateMethod · 0.80

Tested by

no test coverage detected