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

Class PairStr2Pair

rfdiffusion/Track_module.py:72–104  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

70 return msa
71
72class PairStr2Pair(nn.Module):
73 def __init__(self, d_pair=128, n_head=4, d_hidden=32, d_rbf=36, p_drop=0.15):
74 super(PairStr2Pair, self).__init__()
75
76 self.emb_rbf = nn.Linear(d_rbf, d_hidden)
77 self.proj_rbf = nn.Linear(d_hidden, d_pair)
78
79 self.drop_row = Dropout(broadcast_dim=1, p_drop=p_drop)
80 self.drop_col = Dropout(broadcast_dim=2, p_drop=p_drop)
81
82 self.row_attn = BiasedAxialAttention(d_pair, d_pair, n_head, d_hidden, p_drop=p_drop, is_row=True)
83 self.col_attn = BiasedAxialAttention(d_pair, d_pair, n_head, d_hidden, p_drop=p_drop, is_row=False)
84
85 self.ff = FeedForwardLayer(d_pair, 2)
86
87 self.reset_parameter()
88
89 def reset_parameter(self):
90 nn.init.kaiming_normal_(self.emb_rbf.weight, nonlinearity='relu')
91 nn.init.zeros_(self.emb_rbf.bias)
92
93 self.proj_rbf = init_lecun_normal(self.proj_rbf)
94 nn.init.zeros_(self.proj_rbf.bias)
95
96 def forward(self, pair, rbf_feat):
97 B, L = pair.shape[:2]
98
99 rbf_feat = self.proj_rbf(F.relu_(self.emb_rbf(rbf_feat)))
100
101 pair = pair + self.drop_row(self.row_attn(pair, rbf_feat))
102 pair = pair + self.drop_col(self.col_attn(pair, rbf_feat))
103 pair = pair + self.ff(pair)
104 return pair
105
106class MSA2Pair(nn.Module):
107 def __init__(self, d_msa=256, d_pair=128, d_hidden=32, p_drop=0.15):

Callers 2

__init__Method · 0.90
__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected