| 34 | self._load_state_dict(self.model.dist_sigma, rl_state_dict, 'dist_sigma') |
| 35 | |
| 36 | def _load_state_dict(self, il_net, rl_state_dict, key_word): |
| 37 | rl_keys = [k for k in rl_state_dict.keys() if key_word in k] |
| 38 | il_keys = il_net.state_dict().keys() |
| 39 | assert len(rl_keys) == len(il_net.state_dict().keys()), f'mismatch number of layers loading {key_word}' |
| 40 | new_state_dict = OrderedDict() |
| 41 | for k_il, k_rl in zip(il_keys, rl_keys): |
| 42 | new_state_dict[k_il] = rl_state_dict[k_rl] |
| 43 | il_net.load_state_dict(new_state_dict) |
| 44 | |
| 45 | def forward(self, batch): |
| 46 | pass |