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

Class DragPositionNetMultiScale

core/drag_embedding.py:76–151  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

74
75
76class 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)

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected