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

Class RolloutStorage

storage.py:13–281  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

11
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)
42 self.uav_aoi = self.uav_aoi.to(device)
43 self.uav_snr = self.uav_snr.to(device)
44 self.uav_tuse = self.uav_tuse.to(device)
45 self.uav_effort = self.uav_effort.to(device)
46 self.rewards = self.rewards.to(device)
47 self.value_preds = self.value_preds.to(device)
48 self.returns = self.returns.to(device)
49 self.action_log_probs = self.action_log_probs.to(device)
50 self.action_dia = self.action_dia.to(device)
51 self.masks = self.masks.to(device)
52 if params.use_rnn:
53 self.temporal_hidden_states = self.temporal_hidden_states.to(device)
54 if params.use_spatial_att:
55 self.spatial_hidden_states=self.spatial_hidden_states.to(device)
56
57 def insert(self, obs, uav_aoi, uav_snr, uav_tuse, uav_effort, actions, action_log_probs, value_preds, rewards,
58 masks, returns, temporal_hidden_states=None,spatial_hidden_states=None):
59 self.action_dia[self.step].copy_(actions.squeeze())
60 self.action_log_probs[self.step].copy_(action_log_probs.squeeze())
61 self.value_preds[self.step].copy_(value_preds.squeeze())
62 self.rewards[self.step].copy_(rewards.squeeze())
63 self.obs[self.step + 1].copy_(obs.squeeze())
64 self.masks[self.step + 1].copy_(masks.squeeze())
65 self.uav_aoi[self.step + 1].copy_(uav_aoi.squeeze())
66 self.uav_snr[self.step + 1].copy_(uav_snr.squeeze())
67 self.uav_tuse[self.step + 1].copy_(uav_tuse.squeeze())
68 self.uav_effort[self.step + 1].copy_(uav_effort.squeeze())
69
70 # self.next_obs[self.step].copy_(next_obs.squeeze())

Callers 1

trainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected