(self)
| 223 | self.reset_parameter() |
| 224 | |
| 225 | def reset_parameter(self): |
| 226 | # initialize weights to normal distribution |
| 227 | self.embed_x = init_lecun_normal(self.embed_x) |
| 228 | self.embed_e1 = init_lecun_normal(self.embed_e1) |
| 229 | self.embed_e2 = init_lecun_normal(self.embed_e2) |
| 230 | |
| 231 | # initialize bias to zeros |
| 232 | nn.init.zeros_(self.embed_x.bias) |
| 233 | nn.init.zeros_(self.embed_e1.bias) |
| 234 | nn.init.zeros_(self.embed_e2.bias) |
| 235 | |
| 236 | @torch.cuda.amp.autocast(enabled=False) |
| 237 | def forward(self, msa, pair, R_in, T_in, xyz, state, idx, motif_mask, top_k=64, eps=1e-5): |
no test coverage detected