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

Method reset_parameter

rfdiffusion/Track_module.py:159–178  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

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 '''

Callers 1

__init__Method · 0.95

Calls 1

init_lecun_normalFunction · 0.85

Tested by

no test coverage detected