MCPcopy Create free account
hub / github.com/YeWR/EfficientZero / concat_output

Function concat_output

core/model.py:30–50  ·  view source on GitHub ↗
(output_lst)

Source from the content-addressed store, hash-verified

28
29
30def 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
53class BaseNet(nn.Module):

Callers 2

_prepare_reward_valueMethod · 0.90
_prepare_policy_reMethod · 0.90

Calls 1

appendMethod · 0.80

Tested by

no test coverage detected