| 294 | return Ri, Ti, state, alpha |
| 295 | |
| 296 | class IterBlock(nn.Module): |
| 297 | def __init__(self, d_msa=256, d_pair=128, |
| 298 | n_head_msa=8, n_head_pair=4, |
| 299 | use_global_attn=False, |
| 300 | d_hidden=32, d_hidden_msa=None, p_drop=0.15, |
| 301 | SE3_param={'l0_in_features':32, 'l0_out_features':16, 'num_edge_features':32}): |
| 302 | super(IterBlock, self).__init__() |
| 303 | if d_hidden_msa == None: |
| 304 | d_hidden_msa = d_hidden |
| 305 | |
| 306 | self.msa2msa = MSAPairStr2MSA(d_msa=d_msa, d_pair=d_pair, |
| 307 | n_head=n_head_msa, |
| 308 | d_state=SE3_param['l0_out_features'], |
| 309 | use_global_attn=use_global_attn, |
| 310 | d_hidden=d_hidden_msa, p_drop=p_drop) |
| 311 | self.msa2pair = MSA2Pair(d_msa=d_msa, d_pair=d_pair, |
| 312 | d_hidden=d_hidden//2, p_drop=p_drop) |
| 313 | #d_hidden=d_hidden, p_drop=p_drop) |
| 314 | self.pair2pair = PairStr2Pair(d_pair=d_pair, n_head=n_head_pair, |
| 315 | d_hidden=d_hidden, p_drop=p_drop) |
| 316 | self.str2str = Str2Str(d_msa=d_msa, d_pair=d_pair, |
| 317 | d_state=SE3_param['l0_out_features'], |
| 318 | SE3_param=SE3_param, |
| 319 | p_drop=p_drop) |
| 320 | |
| 321 | def forward(self, msa, pair, R_in, T_in, xyz, state, idx, motif_mask, use_checkpoint=False): |
| 322 | rbf_feat = rbf(torch.cdist(xyz[:,:,1,:], xyz[:,:,1,:])) |
| 323 | if use_checkpoint: |
| 324 | msa = checkpoint.checkpoint(create_custom_forward(self.msa2msa), msa, pair, rbf_feat, state) |
| 325 | pair = checkpoint.checkpoint(create_custom_forward(self.msa2pair), msa, pair) |
| 326 | pair = checkpoint.checkpoint(create_custom_forward(self.pair2pair), pair, rbf_feat) |
| 327 | R, T, state, alpha = checkpoint.checkpoint(create_custom_forward(self.str2str, top_k=0), msa, pair, R_in, T_in, xyz, state, idx, motif_mask) |
| 328 | else: |
| 329 | msa = self.msa2msa(msa, pair, rbf_feat, state) |
| 330 | pair = self.msa2pair(msa, pair) |
| 331 | pair = self.pair2pair(pair, rbf_feat) |
| 332 | R, T, state, alpha = self.str2str(msa, pair, R_in, T_in, xyz, state, idx, motif_mask=motif_mask, top_k=0) |
| 333 | |
| 334 | return msa, pair, R, T, state, alpha |
| 335 | |
| 336 | class IterativeSimulator(nn.Module): |
| 337 | def __init__(self, n_extra_block=4, n_main_block=12, n_ref_block=4, |