MCPcopy Create free account
hub / github.com/Sin3DM/Sin3DM / update_network

Method update_network

src/encoding/model.py:178–184  ·  view source on GitHub ↗

update network by back propagation

(self, loss_dict)

Source from the content-addressed store, hash-verified

176 self.net.to(self.device)
177
178 def update_network(self, loss_dict):
179 """update network by back propagation"""
180 loss = sum(loss_dict.values())
181 self.optimizer.zero_grad()
182 loss.backward()
183 self.optimizer.step()
184 self.scheduler.step()
185
186 def _forward_batch(self, data):
187 """forward a batch of data"""

Callers 1

trainMethod · 0.95

Calls 2

zero_gradMethod · 0.80
backwardMethod · 0.45

Tested by

no test coverage detected