MCPcopy Create free account
hub / github.com/ant-research/CoDeF / deform_pts

Method deform_pts

train.py:142–175  ·  view source on GitHub ↗
(self, ts_w, grid, encode_w, step=0, i=0)

Source from the content-addressed store, hash-verified

140 self.models_to_train += [self.models]
141
142 def deform_pts(self, ts_w, grid, encode_w, step=0, i=0):
143 if hparams.deform_hash:
144 ts_w_norm = ts_w / self.seq_len
145 ts_w_norm = ts_w_norm.repeat(grid.shape[0], 1)
146 input_xyt = torch.cat([grid, ts_w_norm], dim=-1)
147 if 'aneal_hash' in self.embeddings.keys():
148 deform = self.models[f'warping_field_{i}'](
149 input_xyt,
150 step=step,
151 aneal_func=self.embeddings['aneal_hash'])
152 else:
153 deform = self.models[f'warping_field_{i}'](input_xyt)
154 if encode_w:
155 deformed_grid = deform + grid
156 else:
157 deformed_grid = grid
158 else:
159 if encode_w:
160 e_w = self.embeddings[f'w_{i}'](repeat(ts_w, 'b n -> (b l) n ',
161 l=grid.shape[0])[:, 0])
162 # Whether to use annealed positional encoding.
163 if self.hparams.annealed:
164 pe_w = self.embeddings['xyz_w'][i](grid, step)
165 else:
166 pe_w = self.embeddings['xyz_w'][i](grid)
167
168 # Warping field type.
169 deform = self.models[f'warping_field_{i}'](torch.cat(
170 [e_w, pe_w], 1))
171 deformed_grid = deform + grid
172 else:
173 deformed_grid = grid
174
175 return deformed_grid
176
177 def forward(self,
178 ts_w,

Callers 1

forwardMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected