| 74 | |
| 75 | |
| 76 | class DragPositionNetMultiScale(nn.Module): |
| 77 | def __init__(self, fourier_freqs=8, scales=[256, 128, 64, 32, 16, 8], channels=[64, 64, 128, 256, 512, 1024], drag_layer_idx=None): |
| 78 | super().__init__() |
| 79 | self.fourier_embedder = FourierEmbedder(num_freqs=fourier_freqs) |
| 80 | self.position_dim = fourier_freqs*2*2 # 2 for sin and cos, 2 for 2 dims (x1, y1) or (x2, y2) |
| 81 | |
| 82 | # -------------------------------------------------------------- # |
| 83 | self.linear_drags = [] |
| 84 | for i in range(len(channels)): |
| 85 | if drag_layer_idx is not None and i != drag_layer_idx: |
| 86 | continue |
| 87 | self.linear_drags.append(nn.Sequential( |
| 88 | nn.Linear(self.position_dim, 128), |
| 89 | nn.SiLU(), |
| 90 | nn.Linear(128, 256), |
| 91 | nn.SiLU(), |
| 92 | nn.Linear(256, channels[i]//2), |
| 93 | )) |
| 94 | |
| 95 | self.linears_drags = nn.ModuleList(self.linear_drags) |
| 96 | self.gaussian_blur = GaussianBlur(kernel_size=5, sigma=1.0) |
| 97 | |
| 98 | self.scales = scales |
| 99 | self.channels = channels |
| 100 | |
| 101 | if drag_layer_idx is not None: |
| 102 | self.drag_layer_idx = drag_layer_idx |
| 103 | self.scales = self.scales[self.drag_layer_idx:self.drag_layer_idx+1] |
| 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) |