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

Class BaseNet

core/model.py:53–131  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

51
52
53class 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())

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected