| 51 | |
| 52 | |
| 53 | class BaseNet(nn.Module): |
| 54 | def __init__(self, inverse_value_transform, inverse_reward_transform, lstm_hidden_size): |
| 55 | """Base Network |
| 56 | schedule_timesteps. After this many timesteps pass final_p is |
| 57 | returned. |
| 58 | Parameters |
| 59 | ---------- |
| 60 | inverse_value_transform: Any |
| 61 | A function that maps value supports into value scalars |
| 62 | inverse_reward_transform: Any |
| 63 | A function that maps reward supports into value scalars |
| 64 | lstm_hidden_size: int |
| 65 | dim of lstm hidden |
| 66 | """ |
| 67 | super(BaseNet, self).__init__() |
| 68 | self.inverse_value_transform = inverse_value_transform |
| 69 | self.inverse_reward_transform = inverse_reward_transform |
| 70 | self.lstm_hidden_size = lstm_hidden_size |
| 71 | |
| 72 | def prediction(self, state): |
| 73 | raise NotImplementedError |
| 74 | |
| 75 | def representation(self, obs_history): |
| 76 | raise NotImplementedError |
| 77 | |
| 78 | def dynamics(self, state, reward_hidden, action): |
| 79 | raise NotImplementedError |
| 80 | |
| 81 | def initial_inference(self, obs) -> NetworkOutput: |
| 82 | num = obs.size(0) |
| 83 | |
| 84 | state = self.representation(obs) |
| 85 | actor_logit, value = self.prediction(state) |
| 86 | |
| 87 | if not self.training: |
| 88 | # if not in training, obtain the scalars of the value/reward |
| 89 | value = self.inverse_value_transform(value).detach().cpu().numpy() |
| 90 | state = state.detach().cpu().numpy() |
| 91 | actor_logit = actor_logit.detach().cpu().numpy() |
| 92 | # zero initialization for reward (value prefix) hidden states |
| 93 | reward_hidden = (torch.zeros(1, num, self.lstm_hidden_size).detach().cpu().numpy(), |
| 94 | torch.zeros(1, num, self.lstm_hidden_size).detach().cpu().numpy()) |
| 95 | else: |
| 96 | # zero initialization for reward (value prefix) hidden states |
| 97 | reward_hidden = (torch.zeros(1, num, self.lstm_hidden_size).to('cuda'), torch.zeros(1, num, self.lstm_hidden_size).to('cuda')) |
| 98 | |
| 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()) |
nothing calls this directly
no outgoing calls
no test coverage detected