(self, inputs, embeddings, bound=1)
| 205 | return f"GridEncoder: input_dim={self.input_dim} num_levels={self.num_levels} level_dim={self.level_dim} resolution={self.base_resolution} -> {int(round(self.base_resolution * self.per_level_scale ** (self.num_levels - 1)))} per_level_scale={self.per_level_scale:.4f} params={tuple(self.embeddings.shape)} gridtype={self.gridtype} align_corners={self.align_corners}" |
| 206 | |
| 207 | def forward(self, inputs, embeddings, bound=1): |
| 208 | # inputs: [..., input_dim], normalized real world positions in [-bound, bound] |
| 209 | # return: [..., num_levels * level_dim] |
| 210 | input_embeddings = torch.cat([embeddings, self.embeddings], dim=0) |
| 211 | |
| 212 | inputs = (inputs + bound) / (2 * bound) # map to [0, 1] |
| 213 | |
| 214 | #print('inputs', inputs.shape, inputs.dtype, inputs.min().item(), inputs.max().item()) |
| 215 | |
| 216 | prefix_shape = list(inputs.shape[:-1]) |
| 217 | inputs = inputs.view(-1, self.input_dim) |
| 218 | |
| 219 | outputs = grid_encode(inputs, input_embeddings, self.offsets, self.per_level_scale, self.base_resolution, inputs.requires_grad, self.gridtype_id, self.align_corners) |
| 220 | outputs = outputs.view(prefix_shape + [self.output_dim]) |
| 221 | |
| 222 | #print('outputs', outputs.shape, outputs.dtype, outputs.min().item(), outputs.max().item()) |
| 223 | |
| 224 | return outputs |
nothing calls this directly
no outgoing calls
no test coverage detected