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

Class SCPred

rfdiffusion/Track_module.py:140–198  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

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):
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))

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected