(self, drags_start, drags_end)
| 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): |
nothing calls this directly
no outgoing calls
no test coverage detected