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

Class Str2Str

rfdiffusion/Track_module.py:201–294  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

199
200
201class Str2Str(nn.Module):
202 def __init__(self, d_msa=256, d_pair=128, d_state=16,
203 SE3_param={'l0_in_features':32, 'l0_out_features':16, 'num_edge_features':32}, p_drop=0.1):
204 super(Str2Str, self).__init__()
205
206 # initial node & pair feature process
207 self.norm_msa = nn.LayerNorm(d_msa)
208 self.norm_pair = nn.LayerNorm(d_pair)
209 self.norm_state = nn.LayerNorm(d_state)
210
211 self.embed_x = nn.Linear(d_msa+d_state, SE3_param['l0_in_features'])
212 self.embed_e1 = nn.Linear(d_pair, SE3_param['num_edge_features'])
213 self.embed_e2 = nn.Linear(SE3_param['num_edge_features']+36+1, SE3_param['num_edge_features'])
214
215 self.norm_node = nn.LayerNorm(SE3_param['l0_in_features'])
216 self.norm_edge1 = nn.LayerNorm(SE3_param['num_edge_features'])
217 self.norm_edge2 = nn.LayerNorm(SE3_param['num_edge_features'])
218
219 self.se3 = SE3TransformerWrapper(**SE3_param)
220 self.sc_predictor = SCPred(d_msa=d_msa, d_state=SE3_param['l0_out_features'],
221 p_drop=p_drop)
222
223 self.reset_parameter()
224
225 def reset_parameter(self):
226 # initialize weights to normal distribution
227 self.embed_x = init_lecun_normal(self.embed_x)
228 self.embed_e1 = init_lecun_normal(self.embed_e1)
229 self.embed_e2 = init_lecun_normal(self.embed_e2)
230
231 # initialize bias to zeros
232 nn.init.zeros_(self.embed_x.bias)
233 nn.init.zeros_(self.embed_e1.bias)
234 nn.init.zeros_(self.embed_e2.bias)
235
236 @torch.cuda.amp.autocast(enabled=False)
237 def forward(self, msa, pair, R_in, T_in, xyz, state, idx, motif_mask, top_k=64, eps=1e-5):
238 B, N, L = msa.shape[:3]
239
240 if motif_mask is None:
241 motif_mask = torch.zeros(L).bool()
242
243 # process msa & pair features
244 node = self.norm_msa(msa[:,0])
245 pair = self.norm_pair(pair)
246 state = self.norm_state(state)
247
248 node = torch.cat((node, state), dim=-1)
249 node = self.norm_node(self.embed_x(node))
250 pair = self.norm_edge1(self.embed_e1(pair))
251
252 neighbor = get_seqsep(idx)
253 rbf_feat = rbf(torch.cdist(xyz[:,:,1], xyz[:,:,1]))
254 pair = torch.cat((pair, rbf_feat, neighbor), dim=-1)
255 pair = self.norm_edge2(self.embed_e2(pair))
256
257 # define graph
258 if top_k != 0:

Callers 2

__init__Method · 0.85
__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected