(self, hidden_state, reward_hidden, action)
| 99 | return NetworkOutput(value, [0. for _ in range(num)], actor_logit, state, reward_hidden) |
| 100 | |
| 101 | def recurrent_inference(self, hidden_state, reward_hidden, action) -> NetworkOutput: |
| 102 | state, reward_hidden, value_prefix = self.dynamics(hidden_state, reward_hidden, action) |
| 103 | actor_logit, value = self.prediction(state) |
| 104 | |
| 105 | if not self.training: |
| 106 | # if not in training, obtain the scalars of the value/reward |
| 107 | value = self.inverse_value_transform(value).detach().cpu().numpy() |
| 108 | value_prefix = self.inverse_reward_transform(value_prefix).detach().cpu().numpy() |
| 109 | state = state.detach().cpu().numpy() |
| 110 | reward_hidden = (reward_hidden[0].detach().cpu().numpy(), reward_hidden[1].detach().cpu().numpy()) |
| 111 | actor_logit = actor_logit.detach().cpu().numpy() |
| 112 | |
| 113 | return NetworkOutput(value, value_prefix, actor_logit, state, reward_hidden) |
| 114 | |
| 115 | def get_weights(self): |
| 116 | return {k: v.cpu() for k, v in self.state_dict().items()} |
no test coverage detected