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

Method recurrent_inference

core/model.py:101–113  ·  view source on GitHub ↗
(self, hidden_state, reward_hidden, action)

Source from the content-addressed store, hash-verified

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()}

Callers 2

update_weightsFunction · 0.80
searchMethod · 0.80

Calls 5

dynamicsMethod · 0.95
predictionMethod · 0.95
NetworkOutputClass · 0.85

Tested by

no test coverage detected