(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)
| 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) |
nothing calls this directly
no outgoing calls
no test coverage detected