MCPcopy Create free account
hub / github.com/RosettaCommons/RFdiffusion / __init__

Method __init__

rfdiffusion/Track_module.py:297–319  ·  view source on GitHub ↗
(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})

Source from the content-addressed store, hash-verified

295
296class 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,:]))

Callers

nothing calls this directly

Calls 5

MSAPairStr2MSAClass · 0.85
MSA2PairClass · 0.85
PairStr2PairClass · 0.85
Str2StrClass · 0.85
__init__Method · 0.45

Tested by

no test coverage detected