(self, drags_start, drags_end)
| 104 | self.channels = self.channels[self.drag_layer_idx:self.drag_layer_idx+1] |
| 105 | |
| 106 | def forward(self, drags_start, drags_end): |
| 107 | scales = self.scales |
| 108 | channels = self.channels |
| 109 | |
| 110 | B, V, N, _ = drags_start.shape |
| 111 | |
| 112 | drags_start = drags_start.view(B*V, N, -1) |
| 113 | drags_end = drags_end.view(B*V, N, -1) |
| 114 | |
| 115 | |
| 116 | multi_scale_merge_start_embeddings = [] |
| 117 | multi_scale_merge_end_embeddings = [] |
| 118 | |
| 119 | for idx, scale in enumerate(scales): |
| 120 | drags_start_embeddings = [] |
| 121 | drags_end_embeddings = [] |
| 122 | for i in range(N): |
| 123 | drag_start_embedding = self.fourier_embedder(drags_start[:, i, :]) |
| 124 | drags_start_embeddings.append(self.linears_drags[idx](drag_start_embedding)) |
| 125 | drags_start_embeddings = torch.stack(drags_start_embeddings, dim=1) |
| 126 | |
| 127 | for i in range(N): |
| 128 | drag_end_embedding = self.fourier_embedder(drags_end[:, i, :]) |
| 129 | drags_end_embeddings.append(self.linears_drags[idx](drag_end_embedding)) |
| 130 | drags_end_embeddings = torch.stack(drags_end_embeddings, dim=1) |
| 131 | |
| 132 | merge_start_embeddings = torch.zeros((B*V, channels[idx]//2, scale, scale)).to(drag_start_embedding.device) |
| 133 | merge_end_embeddings = torch.zeros((B*V, channels[idx]//2, scale, scale)).to(drag_start_embedding.device) |
| 134 | downsample_ratio = 512 // scale |
| 135 | |
| 136 | for i in range(B*V): |
| 137 | for j in range(N): |
| 138 | merge_start_embeddings[i, :, int(drags_start[i, j, 0]) // downsample_ratio, |
| 139 | int(drags_start[i, j, 1]) // downsample_ratio] += drags_start_embeddings[i,j,:] |
| 140 | merge_end_embeddings[i, :, int(drags_end[i, j, 0]) // downsample_ratio, |
| 141 | int(drags_end[i, j, 1]) // downsample_ratio] += drags_end_embeddings[i,j, :] |
| 142 | # merge_end_embeddings[i, :, int(drags_start[i, j, 0]) // downsample_ratio, |
| 143 | # int(drags_start[i, j, 1]) // downsample_ratio] += drags_end_embeddings[i,j, :] |
| 144 | # Add Gaussian Blur |
| 145 | merge_start_embeddings = self.gaussian_blur(merge_start_embeddings) |
| 146 | merge_end_embeddings = self.gaussian_blur(merge_end_embeddings) |
| 147 | |
| 148 | multi_scale_merge_start_embeddings.append(merge_start_embeddings) |
| 149 | multi_scale_merge_end_embeddings.append(merge_end_embeddings) |
| 150 | |
| 151 | return multi_scale_merge_start_embeddings, multi_scale_merge_end_embeddings |
| 152 |
nothing calls this directly
no outgoing calls
no test coverage detected