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

Class DragPositionNet

core/drag_embedding.py:23–73  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

21
22
23class DragPositionNet(nn.Module):
24 def __init__(self, num_drags, fourier_freqs=8, downsample_ratio=64):
25 super().__init__()
26 self.num_drags = num_drags
27
28 self.fourier_embedder = FourierEmbedder(num_freqs=fourier_freqs)
29 self.position_dim = fourier_freqs*2*2 # 2 for sin and cos, 2 for 2 dims (x1, y1) or (x2, y2)
30
31 # -------------------------------------------------------------- #
32 self.linears_drag = nn.Sequential(
33 nn.Linear(self.position_dim, 128),
34 nn.SiLU(),
35 nn.Linear(128, 256),
36 nn.SiLU(),
37 nn.Linear(256, 512),
38 )
39
40 self.downsample_ratio = downsample_ratio
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 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected