| 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, |