(output_lst)
| 28 | |
| 29 | |
| 30 | def concat_output(output_lst): |
| 31 | # concat the model output |
| 32 | value_lst, reward_lst, policy_logits_lst, hidden_state_lst = [], [], [], [] |
| 33 | reward_hidden_c_lst, reward_hidden_h_lst =[], [] |
| 34 | for output in output_lst: |
| 35 | value_lst.append(output.value) |
| 36 | reward_lst.append(output.value_prefix) |
| 37 | policy_logits_lst.append(output.policy_logits) |
| 38 | hidden_state_lst.append(output.hidden_state) |
| 39 | reward_hidden_c_lst.append(output.reward_hidden[0].squeeze(0)) |
| 40 | reward_hidden_h_lst.append(output.reward_hidden[1].squeeze(0)) |
| 41 | |
| 42 | value_lst = np.concatenate(value_lst) |
| 43 | reward_lst = np.concatenate(reward_lst) |
| 44 | policy_logits_lst = np.concatenate(policy_logits_lst) |
| 45 | # hidden_state_lst = torch.cat(hidden_state_lst, 0) |
| 46 | hidden_state_lst = np.concatenate(hidden_state_lst) |
| 47 | reward_hidden_c_lst = np.expand_dims(np.concatenate(reward_hidden_c_lst), axis=0) |
| 48 | reward_hidden_h_lst = np.expand_dims(np.concatenate(reward_hidden_h_lst), axis=0) |
| 49 | |
| 50 | return value_lst, reward_lst, policy_logits_lst, hidden_state_lst, (reward_hidden_c_lst, reward_hidden_h_lst) |
| 51 | |
| 52 | |
| 53 | class BaseNet(nn.Module): |
no test coverage detected