| 138 | return pair |
| 139 | |
| 140 | class SCPred(nn.Module): |
| 141 | def __init__(self, d_msa=256, d_state=32, d_hidden=128, p_drop=0.15): |
| 142 | super(SCPred, self).__init__() |
| 143 | self.norm_s0 = nn.LayerNorm(d_msa) |
| 144 | self.norm_si = nn.LayerNorm(d_state) |
| 145 | self.linear_s0 = nn.Linear(d_msa, d_hidden) |
| 146 | self.linear_si = nn.Linear(d_state, d_hidden) |
| 147 | |
| 148 | # ResNet layers |
| 149 | self.linear_1 = nn.Linear(d_hidden, d_hidden) |
| 150 | self.linear_2 = nn.Linear(d_hidden, d_hidden) |
| 151 | self.linear_3 = nn.Linear(d_hidden, d_hidden) |
| 152 | self.linear_4 = nn.Linear(d_hidden, d_hidden) |
| 153 | |
| 154 | # Final outputs |
| 155 | self.linear_out = nn.Linear(d_hidden, 20) |
| 156 | |
| 157 | self.reset_parameter() |
| 158 | |
| 159 | def reset_parameter(self): |
| 160 | # normal initialization |
| 161 | self.linear_s0 = init_lecun_normal(self.linear_s0) |
| 162 | self.linear_si = init_lecun_normal(self.linear_si) |
| 163 | self.linear_out = init_lecun_normal(self.linear_out) |
| 164 | nn.init.zeros_(self.linear_s0.bias) |
| 165 | nn.init.zeros_(self.linear_si.bias) |
| 166 | nn.init.zeros_(self.linear_out.bias) |
| 167 | |
| 168 | # right before relu activation: He initializer (kaiming normal) |
| 169 | nn.init.kaiming_normal_(self.linear_1.weight, nonlinearity='relu') |
| 170 | nn.init.zeros_(self.linear_1.bias) |
| 171 | nn.init.kaiming_normal_(self.linear_3.weight, nonlinearity='relu') |
| 172 | nn.init.zeros_(self.linear_3.bias) |
| 173 | |
| 174 | # right before residual connection: zero initialize |
| 175 | nn.init.zeros_(self.linear_2.weight) |
| 176 | nn.init.zeros_(self.linear_2.bias) |
| 177 | nn.init.zeros_(self.linear_4.weight) |
| 178 | nn.init.zeros_(self.linear_4.bias) |
| 179 | |
| 180 | def forward(self, seq, state): |
| 181 | ''' |
| 182 | Predict side-chain torsion angles along with backbone torsions |
| 183 | Inputs: |
| 184 | - seq: hidden embeddings corresponding to query sequence (B, L, d_msa) |
| 185 | - state: state feature (output l0 feature) from previous SE3 layer (B, L, d_state) |
| 186 | Outputs: |
| 187 | - si: predicted torsion angles (phi, psi, omega, chi1~4 with cos/sin, Cb bend, Cb twist, CG) (B, L, 10, 2) |
| 188 | ''' |
| 189 | B, L = seq.shape[:2] |
| 190 | seq = self.norm_s0(seq) |
| 191 | state = self.norm_si(state) |
| 192 | si = self.linear_s0(seq) + self.linear_si(state) |
| 193 | |
| 194 | si = si + self.linear_2(F.relu_(self.linear_1(F.relu_(si)))) |
| 195 | si = si + self.linear_4(F.relu_(self.linear_3(F.relu_(si)))) |
| 196 | |
| 197 | si = self.linear_out(F.relu_(si)) |