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

Method apply

examples/cnn_ms/train_ms_model.py:143–183  ·  view source on GitHub ↗

Performs a single optimization step. Args: param_name(String): the name of the param param_value(Tensor): param values to be update in-place grad(Tensor): param gradients; the values may be updated in this function; cann

(self, param_name, param_value, param_grad)

Source from the content-addressed store, hash-verified

141 "Nesterov momentum requires a momentum and zero dampening")
142
143 def apply(self, param_name, param_value, param_grad):
144 """Performs a single optimization step.
145 Args:
146 param_name(String): the name of the param
147 param_value(Tensor): param values to be update in-place
148 grad(Tensor): param gradients; the values may be updated
149 in this function; cannot use it anymore
150 """
151 assert param_value.shape == param_grad.shape, ("shape mismatch",
152 param_value.shape,
153 param_grad.shape)
154 self.device_check(param_value, self.step_counter, self.lr_value,
155 self.mom_value, self.dam_value, self.decay_value)
156
157 # derive dtype from input
158 assert param_value.dtype == self.dtype
159
160 # TODO add branch operator
161 # if self.decay_value != 0:
162 if self.weight_decay.init_value != 0:
163 singa.Axpy(self.decay_value.data, param_value.data, param_grad.data)
164
165 if self.momentum.init_value != 0:
166 if param_name not in self.moments:
167 flag = param_value.device.graph_enabled()
168 param_value.device.EnableGraph(False)
169 self.moments[param_name] = tensor.zeros_like(param_value)
170 param_value.device.EnableGraph(flag)
171
172 buf = self.moments[param_name]
173 buf *= self.mom_value
174 alpha = 1.0 - self.dam_value
175 singa.Axpy(alpha.data, param_grad.data, buf.data)
176
177 if self.nesterov:
178 singa.Axpy(self.mom_value.data, buf.data, param_grad.data)
179 else:
180 param_grad = buf
181
182 minus_lr = 0.0 - self.lr_value
183 singa.Axpy(minus_lr.data, param_grad.data, param_value.data)
184
185 def step(self):
186 # increment step counter, lr and moment

Callers 1

call_with_returnsMethod · 0.45

Calls 3

graph_enabledMethod · 0.80
EnableGraphMethod · 0.80
device_checkMethod · 0.45

Tested by

no test coverage detected