(multires, i=0)
| 234 | |
| 235 | |
| 236 | def get_embedder(multires, i=0): |
| 237 | import torch.nn as nn |
| 238 | if i == -1: |
| 239 | return nn.Identity(), 3 |
| 240 | |
| 241 | embed_kwargs = { |
| 242 | 'include_input': True, |
| 243 | 'input_dims': 3, |
| 244 | 'max_freq_log2': multires - 1, |
| 245 | 'num_freqs': multires, |
| 246 | 'log_sampling': True, |
| 247 | 'periodic_fns': [torch.sin, torch.cos], |
| 248 | } |
| 249 | |
| 250 | embedder_obj = Embedder(**embed_kwargs) |
| 251 | embed = lambda x, eo=embedder_obj: eo.embed(x) |
| 252 | return embed, embedder_obj.out_dim |
| 253 | |
| 254 | |
| 255 | class APOPMeter(): |