MCPcopy Create free account
hub / github.com/BIT-MCS/DRL-eFresh / __init__

Method __init__

storage.py:14–38  ·  view source on GitHub ↗
(self, num_steps, mini_batch_num, obs_shape, uav_num,
                 temporal_hidden_state_size=params.temporal_hidden_size,
                 spatial_hidden_state_size=params.spatial_hidden_size)

Source from the content-addressed store, hash-verified

12
13class RolloutStorage(object):
14 def __init__(self, num_steps, mini_batch_num, obs_shape, uav_num,
15 temporal_hidden_state_size=params.temporal_hidden_size,
16 spatial_hidden_state_size=params.spatial_hidden_size):
17 self.mini_batch_num = mini_batch_num
18 self.obs = torch.zeros(num_steps + 1, *obs_shape)
19 self.uav_aoi = torch.zeros(num_steps + 1, params.uav_num)
20 self.uav_snr = torch.zeros(num_steps + 1, params.uav_num)
21 self.uav_tuse = torch.zeros(num_steps + 1, params.uav_num)
22 self.uav_effort = torch.zeros(num_steps + 1, params.uav_num)
23 # self.next_obs = torch.zeros(num_steps + 1, *obs_shape)
24 self.rewards = torch.zeros(num_steps, 1)
25
26 self.value_preds = torch.zeros(num_steps + 1, 1)
27 self.returns = torch.zeros(num_steps + 1, 1)
28 self.action_log_probs = torch.zeros(num_steps, 1)
29 self.action_dia = torch.zeros(num_steps, int(uav_num * params.uav_action_dim))
30 self.masks = torch.ones(num_steps + 1, 1)
31 self.num_steps = num_steps
32 self.step = 0
33 if params.use_rnn:
34 self.temporal_hidden_states = torch.zeros(num_steps + 1, temporal_hidden_state_size)
35 self.temporal_hidden_state_size = temporal_hidden_state_size
36 if params.use_spatial_att:
37 self.spatial_hidden_states=torch.zeros(num_steps + 1, spatial_hidden_state_size,8, 8)
38 self.spatial_hidden_state_size = spatial_hidden_state_size
39
40 def to(self, device):
41 self.obs = self.obs.to(device)

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected