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

Method forward

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

Source from the content-addressed store, hash-verified

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
76class DragPositionNetMultiScale(nn.Module):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected