| 156 | return outputs |
| 157 | |
| 158 | class VarGridEncoder(nn.Module): |
| 159 | def __init__(self, input_dim=3, num_levels=16, level_dim=2, per_level_scale=2, base_resolution=16, log2_hashmap_size=19, desired_resolution=None, gridtype='hash', align_corners=False, hash_entries=None): |
| 160 | super().__init__() |
| 161 | |
| 162 | # the finest resolution desired at the last level, if provided, overridee per_level_scale |
| 163 | if desired_resolution is not None: |
| 164 | per_level_scale = np.exp2(np.log2(desired_resolution / base_resolution) / (num_levels - 1)) |
| 165 | |
| 166 | self.input_dim = input_dim # coord dims, 2 or 3 |
| 167 | self.num_levels = num_levels # num levels, each level multiply resolution by 2 |
| 168 | self.level_dim = level_dim # encode channels per level |
| 169 | self.per_level_scale = per_level_scale # multiply resolution by this scale at each level. |
| 170 | self.log2_hashmap_size = log2_hashmap_size |
| 171 | self.base_resolution = base_resolution |
| 172 | self.output_dim = num_levels * level_dim |
| 173 | self.gridtype = gridtype |
| 174 | self.gridtype_id = _gridtype_to_id[gridtype] # "tiled" or "hash" |
| 175 | self.align_corners = align_corners |
| 176 | |
| 177 | # allocate parameters |
| 178 | offsets = [] |
| 179 | offset = 0 |
| 180 | self.max_params = 2 ** log2_hashmap_size |
| 181 | for i in range(num_levels): |
| 182 | resolution = int(np.ceil(base_resolution * per_level_scale ** i)) |
| 183 | params_in_level = min(self.max_params, (resolution if align_corners else resolution + 1) ** input_dim) # limit max number |
| 184 | params_in_level = int(np.ceil(params_in_level / 8) * 8) # make divisible |
| 185 | offsets.append(offset) |
| 186 | offset += params_in_level |
| 187 | offsets.append(offset) |
| 188 | offsets = torch.from_numpy(np.array(offsets, dtype=np.int32)) |
| 189 | self.register_buffer('offsets', offsets) |
| 190 | |
| 191 | self.n_params = offsets[-1] * level_dim |
| 192 | self.level_dim = level_dim |
| 193 | self.offset = offset |
| 194 | |
| 195 | # parameters |
| 196 | self.embeddings = nn.Parameter(torch.empty(offset - hash_entries, level_dim)) |
| 197 | |
| 198 | self.reset_parameters() |
| 199 | |
| 200 | def reset_parameters(self): |
| 201 | std = 1e-4 |
| 202 | self.embeddings.data.uniform_(-std, std) |
| 203 | |
| 204 | def __repr__(self): |
| 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 | |