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

Method train_one_batch

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

Source from the content-addressed store, hash-verified

105 return y
106
107 def train_one_batch(self, x, y, dist_option, spars, synflow_flag):
108 # print ("in train_one_batch")
109 out = self.forward(x)
110 # print ("train_one_batch x.data: \n", x.data)
111 # print ("train_one_batch y.data: \n", y.data)
112 # print ("train_one_batch out.data: \n", out.data)
113 if synflow_flag:
114 # print ("sum_error")
115 loss = self.sum_error(out)
116 else: # normal training
117 # print ("softmax_cross_entropy")
118 loss = self.softmax_cross_entropy(out, y)
119 # print ("train_one_batch loss.data: \n", loss.data)
120
121 if dist_option == 'plain':
122 # print ("before pn_p_g_list = self.optimizer(loss)")
123 pn_p_g_list = self.optimizer(loss)
124 # print ("after pn_p_g_list = self.optimizer(loss)")
125 elif dist_option == 'half':
126 self.optimizer.backward_and_update_half(loss)
127 elif dist_option == 'partialUpdate':
128 self.optimizer.backward_and_partial_update(loss)
129 elif dist_option == 'sparseTopK':
130 self.optimizer.backward_and_sparse_update(loss,
131 topK=True,
132 spars=spars)
133 elif dist_option == 'sparseThreshold':
134 self.optimizer.backward_and_sparse_update(loss,
135 topK=False,
136 spars=spars)
137 # print ("len(pn_p_g_list): \n", len(pn_p_g_list))
138 # print ("len(pn_p_g_list[0]): \n", len(pn_p_g_list[0]))
139 # print ("pn_p_g_list[0][0]: \n", pn_p_g_list[0][0])
140 # print ("pn_p_g_list[0][1].data: \n", pn_p_g_list[0][1].data)
141 # print ("pn_p_g_list[0][2].data: \n", pn_p_g_list[0][2].data)
142 return pn_p_g_list, out, loss
143 # return pn_p_g_list[0], pn_p_g_list[1], pn_p_g_list[2], out, loss
144
145 def set_optimizer(self, optimizer):
146 self.optimizer = optimizer

Callers

nothing calls this directly

Calls 4

forwardMethod · 0.95

Tested by

no test coverage detected