(self, model_input)
| 794 | print(self) |
| 795 | |
| 796 | def forward(self, model_input): |
| 797 | |
| 798 | input_dict = {key: input.clone().detach().requires_grad_(True) |
| 799 | for key, input in model_input.items()} |
| 800 | |
| 801 | if self.nl != 'sine': |
| 802 | for input_to_encode, input_dim, num_pe_fns in self.input_pe_params: |
| 803 | encoded_input = self.positional_encoding_fn[input_to_encode](input_dict[input_to_encode]) |
| 804 | input_dict.update({input_to_encode: encoded_input}) |
| 805 | |
| 806 | input_list = [] |
| 807 | for input_name, _, _ in self.input_pe_params: |
| 808 | input_list.append(input_dict[input_name]) |
| 809 | |
| 810 | coords = torch.cat(input_list, dim=-1) |
| 811 | |
| 812 | if coords.ndim == 2: |
| 813 | coords = coords[None, :, :] |
| 814 | |
| 815 | output = self.net(coords) |
| 816 | |
| 817 | output_dict = {'output': output} |
| 818 | return {'model_in': input_dict, 'model_out': output_dict} |
nothing calls this directly
no outgoing calls
no test coverage detected