| 199 | |
| 200 | |
| 201 | class 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: |