MCPcopy Create free account
hub / github.com/pytorch/pytorch / _run

Method _run

caffe2/python/optimizer.py:1432–1474  ·  view source on GitHub ↗
(self, net, param_init_net, param_info)

Source from the content-addressed store, hash-verified

1430 self.init_kwargs = kwargs
1431
1432 def _run(self, net, param_init_net, param_info):
1433 param = param_info.blob
1434 grad = param_info.grad
1435
1436 if self.alpha <= 0:
1437 return
1438
1439 lr, _ = self.build_lr(
1440 net,
1441 param_init_net,
1442 base_learning_rate=self.alpha,
1443 policy=self.policy,
1444 **(self.init_kwargs)
1445 )
1446
1447 moment = param_init_net.ConstantFill(
1448 [param], str(param) + "_squared_moment", value=0.0
1449 )
1450
1451 moment_update = param_init_net.ConstantFill(
1452 [param], str(param) + "_squared_moment_update", value=0.0
1453 )
1454
1455 self._aux_params.local.append(moment)
1456 self._aux_params.local.append(moment_update)
1457
1458 if isinstance(grad, core.GradientSlice):
1459 grad = self.dedup(net, self.sparse_dedup_aggregator, grad)
1460 net.SparseAdadelta(
1461 [param, moment, moment_update, grad.indices, grad.values, lr],
1462 [param, moment, moment_update],
1463 epsilon=self.epsilon,
1464 decay=self.decay,
1465 engine=self.engine,
1466 )
1467 else:
1468 net.Adadelta(
1469 [param, moment, moment_update, grad, lr],
1470 [param, moment, moment_update],
1471 epsilon=self.epsilon,
1472 decay=self.decay,
1473 engine=self.engine,
1474 )
1475
1476 def scale_learning_rate(self, scale):
1477 self.alpha *= scale

Callers

nothing calls this directly

Calls 4

isinstanceFunction · 0.85
build_lrMethod · 0.80
dedupMethod · 0.80
appendMethod · 0.45

Tested by

no test coverage detected