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

Method __init__

methods/model.py:27–47  ·  view source on GitHub ↗
(self, obs_shape, num_of_uav, device, trainable=True, hidden_size=params.temporal_hidden_size)

Source from the content-addressed store, hash-verified

25
26class 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):

Callers 1

__init__Method · 0.45

Calls 3

DiagGaussianClass · 0.90
NNBaseClass · 0.85
initFunction · 0.70

Tested by

no test coverage detected