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

Class NNBase

methods/model.py:116–302  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

114
115
116class NNBase(nn.Module):
117 def __init__(self, num_inputs, device, trainable=True, hidden_size=params.temporal_hidden_size):
118 super(NNBase, self).__init__()
119 self._feature_size = hidden_size
120
121 if params.use_rnn is False or params.use_spatial_att is False:
122 init_ = lambda m: init(m, nn.init.orthogonal_, lambda x: nn.init.constant_(x, 0),
123 nn.init.calculate_gain('relu'))
124
125 self.feature = nn.Sequential(
126 # input: 3*80*80
127 init_(nn.Conv2d(num_inputs, 32, 8, stride=4)),
128 nn.LayerNorm([32, 19, 19]),
129 nn.ReLU(inplace=True),
130 # input: 32*19*19
131 init_(nn.Conv2d(32, 64, 4, stride=2)),
132 nn.LayerNorm([64, 8, 8]),
133 nn.ReLU(inplace=True),
134 # input: 64*8*8
135 init_(nn.Conv2d(64, 32, 3, stride=1)),
136 nn.LayerNorm([32, 6, 6]),
137 nn.ReLU(inplace=True),
138 # output: 32*6*6
139 ).to(device)
140
141 init_ = lambda m: init(m,
142 nn.init.orthogonal_,
143 lambda x: nn.init.constant_(x, 0))
144
145 self.conv_to_flat = nn.Sequential(
146 Flatten(),
147 init_(nn.Linear(32 * 6 * 6, self._feature_size)),
148 nn.LayerNorm([self._feature_size]),
149 nn.ReLU(inplace=True),
150 ).to(device)
151
152 if params.use_rnn:
153 if params.use_relational_att:
154 # self.gru = RelationalGRU(input_size=hidden_size, hidden_dim=hidden_size, use_att=params.use_att).to(
155 # device)
156 self.gru = RelationalGRU(input_size=hidden_size, hidden_dim=hidden_size, use_att=False).to(
157 device)
158 elif params.use_spatial_att:
159 self.gru = SpatialAttGRU(input_dim=num_inputs).to(device)
160 else:
161 self.gru = nn.GRU(hidden_size, hidden_size).to(device)
162
163 for name, param in self.gru.named_parameters():
164 if 'bias' in name:
165 nn.init.constant_(param, 0)
166 elif 'weight' in name:
167 nn.init.orthogonal_(param)
168
169 if trainable:
170 self.train()
171 else:
172 self.eval()
173 for p in self.parameters():

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected