(self)
| 157 | self.reset_parameter() |
| 158 | |
| 159 | def reset_parameter(self): |
| 160 | # normal initialization |
| 161 | self.linear_s0 = init_lecun_normal(self.linear_s0) |
| 162 | self.linear_si = init_lecun_normal(self.linear_si) |
| 163 | self.linear_out = init_lecun_normal(self.linear_out) |
| 164 | nn.init.zeros_(self.linear_s0.bias) |
| 165 | nn.init.zeros_(self.linear_si.bias) |
| 166 | nn.init.zeros_(self.linear_out.bias) |
| 167 | |
| 168 | # right before relu activation: He initializer (kaiming normal) |
| 169 | nn.init.kaiming_normal_(self.linear_1.weight, nonlinearity='relu') |
| 170 | nn.init.zeros_(self.linear_1.bias) |
| 171 | nn.init.kaiming_normal_(self.linear_3.weight, nonlinearity='relu') |
| 172 | nn.init.zeros_(self.linear_3.bias) |
| 173 | |
| 174 | # right before residual connection: zero initialize |
| 175 | nn.init.zeros_(self.linear_2.weight) |
| 176 | nn.init.zeros_(self.linear_2.bias) |
| 177 | nn.init.zeros_(self.linear_4.weight) |
| 178 | nn.init.zeros_(self.linear_4.bias) |
| 179 | |
| 180 | def forward(self, seq, state): |
| 181 | ''' |
no test coverage detected