(self, obs_shape, num_of_uav, device, trainable=True, hidden_size=params.temporal_hidden_size)
| 25 | |
| 26 | class Model(nn.Module): |
| 27 | def __init__(self, obs_shape, num_of_uav, device, trainable=True, hidden_size=params.temporal_hidden_size): |
| 28 | # todo 1: add parameter trainable=True |
| 29 | super(Model, self).__init__() |
| 30 | # feature extract |
| 31 | self.base = NNBase(obs_shape[0], device, trainable) |
| 32 | # actor |
| 33 | self.dist_dia = DiagGaussian(hidden_size + num_of_uav * 4, params.uav_action_dim * num_of_uav, |
| 34 | device) # continuous, (dx, dy,v,coll_t) |
| 35 | init_ = lambda m: init(m, |
| 36 | nn.init.orthogonal_, |
| 37 | lambda x: nn.init.constant_(x, 0)) |
| 38 | # critic |
| 39 | self.critic = nn.Sequential( |
| 40 | init_(nn.Linear(hidden_size + num_of_uav * 4, 1)) |
| 41 | ) |
| 42 | self.device = device |
| 43 | # todo 2: distinguish train and eval |
| 44 | if trainable: |
| 45 | self.train() |
| 46 | else: |
| 47 | self.eval() |
| 48 | |
| 49 | def act(self, inputs, uav_aoi, uav_snr, uav_compl, uav_tc_compl, temporal_hidden_state=None, mask=None, |
| 50 | spatial_hidden_state=None): |
no test coverage detected