| 103 | self.create_embedding_fn() |
| 104 | |
| 105 | def create_embedding_fn(self): |
| 106 | embed_fns = [] |
| 107 | d = self.kwargs['input_dims'] |
| 108 | out_dim = 0 |
| 109 | if self.kwargs['include_input']: |
| 110 | embed_fns.append(lambda x : x) |
| 111 | out_dim += d |
| 112 | |
| 113 | max_freq = self.kwargs['max_freq_log2'] |
| 114 | self.N_freqs = self.kwargs['num_freqs'] |
| 115 | |
| 116 | if self.kwargs['log_sampling']: |
| 117 | freq_bands = 2.**torch.linspace(0., max_freq, steps=self.N_freqs) # tensor([ 1., 2., 4., 8., 16., 32., 64., 128., 256., 512.]) |
| 118 | else: |
| 119 | freq_bands = torch.linspace(2.**0., 2.**max_freq, steps=self.N_freqs) |
| 120 | |
| 121 | for freq in freq_bands: # 10 iters for 3D location, 4 iters for 2D direction |
| 122 | for p_fn in self.kwargs['periodic_fns']: |
| 123 | embed_fns.append(lambda x, p_fn=p_fn, freq=freq : p_fn(x * freq)) |
| 124 | out_dim += d |
| 125 | self.embed_fns = embed_fns |
| 126 | self.out_dim = out_dim |
| 127 | |
| 128 | def embed(self, inputs): |
| 129 | if self.kwargs['max_freq_log2'] != 0: |