MCPcopy Create free account
hub / github.com/apache/singa / train_one_batch

Method train_one_batch

examples/msmlp/model.py:103–124  ·  view source on GitHub ↗
(self, x, y, synflow_flag, dist_option, spars)

Source from the content-addressed store, hash-verified

101 return y
102
103 def train_one_batch(self, x, y, synflow_flag, dist_option, spars):
104 out = self.forward(x)
105 if synflow_flag:
106 loss = self.sum_error(out)
107 else: # normal training
108 loss = self.softmax_cross_entropy(out, y)
109
110 if dist_option == 'plain':
111 pn_p_g_list = self.optimizer(loss)
112 elif dist_option == 'half':
113 self.optimizer.backward_and_update_half(loss)
114 elif dist_option == 'partialUpdate':
115 self.optimizer.backward_and_partial_update(loss)
116 elif dist_option == 'sparseTopK':
117 self.optimizer.backward_and_sparse_update(loss,
118 topK=True,
119 spars=spars)
120 elif dist_option == 'sparseThreshold':
121 self.optimizer.backward_and_sparse_update(loss,
122 topK=False,
123 spars=spars)
124 return pn_p_g_list, out, loss
125
126 def set_optimizer(self, optimizer):
127 self.optimizer = optimizer

Callers

nothing calls this directly

Calls 4

forwardMethod · 0.95

Tested by

no test coverage detected