MCPcopy Create free account
hub / github.com/ashawkey/RAD-NeRF / GridEncoder

Class GridEncoder

gridencoder/grid.py:96–185  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

94
95
96class GridEncoder(nn.Module):
97 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, interpolation='linear'):
98 super().__init__()
99
100 # the finest resolution desired at the last level, if provided, overridee per_level_scale
101 if desired_resolution is not None:
102 per_level_scale = np.exp2(np.log2(desired_resolution / base_resolution) / (num_levels - 1))
103
104 self.input_dim = input_dim # coord dims, 2 or 3
105 self.num_levels = num_levels # num levels, each level multiply resolution by 2
106 self.level_dim = level_dim # encode channels per level
107 self.per_level_scale = per_level_scale # multiply resolution by this scale at each level.
108 self.log2_hashmap_size = log2_hashmap_size
109 self.base_resolution = base_resolution
110 self.output_dim = num_levels * level_dim
111 self.gridtype = gridtype
112 self.gridtype_id = _gridtype_to_id[gridtype] # "tiled" or "hash"
113 self.interpolation = interpolation
114 self.interp_id = _interp_to_id[interpolation] # "linear" or "smoothstep"
115 self.align_corners = align_corners
116
117 # allocate parameters
118 offsets = []
119 offset = 0
120 self.max_params = 2 ** log2_hashmap_size
121 for i in range(num_levels):
122 resolution = int(np.ceil(base_resolution * per_level_scale ** i))
123 params_in_level = min(self.max_params, (resolution if align_corners else resolution + 1) ** input_dim) # limit max number
124 params_in_level = int(np.ceil(params_in_level / 8) * 8) # make divisible
125 offsets.append(offset)
126 offset += params_in_level
127 offsets.append(offset)
128 offsets = torch.from_numpy(np.array(offsets, dtype=np.int32))
129 self.register_buffer('offsets', offsets)
130
131 self.n_params = offsets[-1] * level_dim
132
133 # parameters
134 self.embeddings = nn.Parameter(torch.empty(offset, level_dim))
135
136 self.reset_parameters()
137
138 def reset_parameters(self):
139 std = 1e-4
140 self.embeddings.data.uniform_(-std, std)
141
142 def __repr__(self):
143 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} interpolation={self.interpolation}"
144
145 def forward(self, inputs, bound=1):
146 # inputs: [..., input_dim], normalized real world positions in [-bound, bound]
147 # return: [..., num_levels * level_dim]
148
149 inputs = (inputs + bound) / (2 * bound) # map to [0, 1]
150
151 #print('inputs', inputs.shape, inputs.dtype, inputs.min().item(), inputs.max().item())
152
153 prefix_shape = list(inputs.shape[:-1])

Callers 1

get_encoderFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected