| 21 | |
| 22 | |
| 23 | class DragPositionNet(nn.Module): |
| 24 | def __init__(self, num_drags, fourier_freqs=8, downsample_ratio=64): |
| 25 | super().__init__() |
| 26 | self.num_drags = num_drags |
| 27 | |
| 28 | self.fourier_embedder = FourierEmbedder(num_freqs=fourier_freqs) |
| 29 | self.position_dim = fourier_freqs*2*2 # 2 for sin and cos, 2 for 2 dims (x1, y1) or (x2, y2) |
| 30 | |
| 31 | # -------------------------------------------------------------- # |
| 32 | self.linears_drag = nn.Sequential( |
| 33 | nn.Linear(self.position_dim, 128), |
| 34 | nn.SiLU(), |
| 35 | nn.Linear(128, 256), |
| 36 | nn.SiLU(), |
| 37 | nn.Linear(256, 512), |
| 38 | ) |
| 39 | |
| 40 | self.downsample_ratio = downsample_ratio |
| 41 | |
| 42 | |
| 43 | def forward(self, drags_start, drags_end): |
| 44 | # drags_start: [B, V, N, 2], start points of drags |
| 45 | # drags_end: [B, V, N, 2], move vectors of drags |
| 46 | B, V, N, _ = drags_start.shape |
| 47 | drags_start = drags_start.view(B*V, N, -1) |
| 48 | |
| 49 | drags_start_embeddings = [] |
| 50 | for i in range(N): |
| 51 | drag_start_embedding = self.fourier_embedder(drags_start[:, i, :]) |
| 52 | drags_start_embeddings.append(self.linears_drag(drag_start_embedding)) |
| 53 | drags_start_embeddings = torch.stack(drags_start_embeddings, dim=1) |
| 54 | |
| 55 | drags_end = drags_end.view(B*V, N, -1) |
| 56 | drags_end_embeddings = [] |
| 57 | for i in range(N): |
| 58 | drag_end_embedding = self.fourier_embedder(drags_end[:, i, :]) |
| 59 | drags_end_embeddings.append(self.linears_drag(drag_end_embedding)) |
| 60 | drags_end_embeddings = torch.stack(drags_end_embeddings, dim=1) |
| 61 | |
| 62 | merge_start_embeddings = torch.zeros((B*V, 512, 8, 8)).to(drag_start_embedding.device) # [B*V, 256, 8, 8] |
| 63 | merge_end_embeddings = torch.zeros((B*V, 512, 8, 8)).to(drag_start_embedding.device) # [B*V, 256, 8, 8] |
| 64 | |
| 65 | for i in range(B*V): |
| 66 | for j in range(N): |
| 67 | merge_start_embeddings[i, :, int(drags_start[i, j, 0]) // self.downsample_ratio, |
| 68 | int(drags_start[i, j, 1]) // self.downsample_ratio] += drags_start_embeddings[i,j,:] |
| 69 | merge_end_embeddings[i, :, int(drags_end[i, j, 0]) // self.downsample_ratio, |
| 70 | int(drags_end[i, j, 1]) // self.downsample_ratio] += drags_end_embeddings[i,j, :] |
| 71 | |
| 72 | merge_embeddings = torch.cat([merge_start_embeddings, merge_end_embeddings], dim=1) |
| 73 | return merge_embeddings |
| 74 | |
| 75 | |
| 76 | class DragPositionNetMultiScale(nn.Module): |