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

Class MSA2Pair

rfdiffusion/Track_module.py:106–138  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

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):
108 super(MSA2Pair, self).__init__()
109 self.norm = nn.LayerNorm(d_msa)
110 self.proj_left = nn.Linear(d_msa, d_hidden)
111 self.proj_right = nn.Linear(d_msa, d_hidden)
112 self.proj_out = nn.Linear(d_hidden*d_hidden, d_pair)
113
114 self.reset_parameter()
115
116 def reset_parameter(self):
117 # normal initialization
118 self.proj_left = init_lecun_normal(self.proj_left)
119 self.proj_right = init_lecun_normal(self.proj_right)
120 nn.init.zeros_(self.proj_left.bias)
121 nn.init.zeros_(self.proj_right.bias)
122
123 # zero initialize output
124 nn.init.zeros_(self.proj_out.weight)
125 nn.init.zeros_(self.proj_out.bias)
126
127 def forward(self, msa, pair):
128 B, N, L = msa.shape[:3]
129 msa = self.norm(msa)
130 left = self.proj_left(msa)
131 right = self.proj_right(msa)
132 right = right / float(N)
133 out = einsum('bsli,bsmj->blmij', left, right).reshape(B, L, L, -1)
134 out = self.proj_out(out)
135
136 pair = pair + out
137
138 return pair
139
140class SCPred(nn.Module):
141 def __init__(self, d_msa=256, d_state=32, d_hidden=128, p_drop=0.15):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected