| 11 | |
| 12 | |
| 13 | class 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()) |