(self, d_msa=256, d_pair=128,
n_head_msa=8, n_head_pair=4,
use_global_attn=False,
d_hidden=32, d_hidden_msa=None, p_drop=0.15,
SE3_param={'l0_in_features':32, 'l0_out_features':16, 'num_edge_features':32})
| 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,:])) |
nothing calls this directly
no test coverage detected