MCPcopy Create free account
hub / github.com/GasaiYU/PartRM / forward

Method forward

core/drag_embedding.py:106–151  ·  view source on GitHub ↗
(self, drags_start, drags_end)

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected